From 6e69abbe954cf2b206d1ff798018135bea341289 Mon Sep 17 00:00:00 2001 From: IceCodeNew <32576256+IceCodeNew@users.noreply.github.com> Date: Sun, 26 Jul 2026 02:32:56 +0800 Subject: [PATCH 1/2] fix(llm): use JSON object structured output --- docs/design.md | 2 +- tests/test_any_llm_compatibility.py | 24 ++- tests/test_any_llm_provider.py | 204 +++++++++++++----- tests/test_cli.py | 6 +- tests/test_config.py | 31 ++- tests/test_llm.py | 2 +- weather_briefing/config/settings.py | 14 +- .../data/any_llm_compatibility.py | 23 ++ weather_briefing/llm/any_llm.py | 52 ++++- 9 files changed, 288 insertions(+), 70 deletions(-) diff --git a/docs/design.md b/docs/design.md index 89505d0f..49315862 100644 --- a/docs/design.md +++ b/docs/design.md @@ -177,7 +177,7 @@ Open-Meteo 的逐小时空气质量和花粉预报按目标日峰值生成生活 开发环境安装 `any-llm-sdk[all]`,用于验证所有 completion provider 的装载边界。基础运行依赖只包含 SDK 核心包。官方镜像额外安装 DeepSeek、OpenAI 和 OpenRouter 所需组件。 -`LLMStructuredOutput` 同时用于 SDK 的结构化输出和应用侧复验。应用还会检查来源 ID、必填建议、预警 ID 和章节间重复等领域规则。 +所有受支持 provider 统一请求 `json_object`,并把 Pydantic JSON Schema 加入最后一条用户消息,避免 OpenAI-compatible 端点只实现 JSON Mode 而拒绝 OpenAI `json_schema`。返回后再用同一 Pydantic 模型严格复验;应用还会检查来源 ID、必填建议、预警 ID 和章节间重复等领域规则。兼容性数据维护与锁定 any-llm SDK 对齐的不支持 JSON Object provider 黑名单,配置入口与 adapter factory 都拒绝黑名单内 provider,不为其他请求格式增加独立分支。 `LLM_MAX_ATTEMPTS` 只修复已经返回但不符合输出契约的正文。认证失败、限流、超时或空响应不进入契约修复。 diff --git a/tests/test_any_llm_compatibility.py b/tests/test_any_llm_compatibility.py index 35245991..0372cf6a 100644 --- a/tests/test_any_llm_compatibility.py +++ b/tests/test_any_llm_compatibility.py @@ -1,6 +1,10 @@ from any_llm import AnyLLM +from any_llm.providers.openai.base import BaseOpenAIProvider -from weather_briefing.data.any_llm_compatibility import UNSUPPORTED_DEFAULT_HEADER_PROVIDERS +from weather_briefing.data.any_llm_compatibility import ( + UNSUPPORTED_DEFAULT_HEADER_PROVIDERS, + UNSUPPORTED_JSON_OBJECT_PROVIDERS, +) def test_default_header_provider_compatibility_matches_the_pinned_sdk() -> None: @@ -25,3 +29,21 @@ def test_default_header_provider_compatibility_matches_the_pinned_sdk() -> None: } assert completion_providers > UNSUPPORTED_DEFAULT_HEADER_PROVIDERS + + +def test_json_object_provider_compatibility_matches_the_pinned_sdk() -> None: + completion_providers = { + provider + for provider in AnyLLM.get_supported_providers() + if AnyLLM.get_provider_class(provider).SUPPORTS_COMPLETION + } + json_object_providers = { + provider + for provider in completion_providers + if issubclass(AnyLLM.get_provider_class(provider), BaseOpenAIProvider) + } + unsupported_providers = completion_providers - json_object_providers + + assert unsupported_providers == UNSUPPORTED_JSON_OBJECT_PROVIDERS + assert json_object_providers | UNSUPPORTED_JSON_OBJECT_PROVIDERS == completion_providers + assert json_object_providers.isdisjoint(UNSUPPORTED_JSON_OBJECT_PROVIDERS) diff --git a/tests/test_any_llm_provider.py b/tests/test_any_llm_provider.py index 859b67b6..f14157c6 100644 --- a/tests/test_any_llm_provider.py +++ b/tests/test_any_llm_provider.py @@ -3,6 +3,7 @@ import os from collections.abc import Callable, Mapping from types import SimpleNamespace +from typing import TypedDict from unittest.mock import AsyncMock, Mock import httpx @@ -14,6 +15,7 @@ from pydantic import BaseModel, ValidationError from weather_briefing.api_client import LoggedAsyncClient +from weather_briefing.data.any_llm_compatibility import UNSUPPORTED_JSON_OBJECT_PROVIDERS from weather_briefing.llm import ( AnyLLMStructuredProvider, FallbackLLMProvider, @@ -26,17 +28,25 @@ from weather_briefing.notifications import NotificationDecision +class _CompletionCall(TypedDict): + model: str + messages: list[dict[str, str]] + response_format: type[BaseModel] | dict[str, object] + temperature: float + max_tokens: int + + class _CompletionClientStub: def __init__(self, response: object) -> None: self._response = response - self.calls: list[dict[str, object]] = [] + self.calls: list[_CompletionCall] = [] async def acompletion( self, *, model: str, messages: list[dict[str, str]], - response_format: type[BaseModel], + response_format: type[BaseModel] | dict[str, object], temperature: float, max_tokens: int, ) -> object: @@ -97,7 +107,7 @@ async def test_service_status_llm_is_created_only_on_first_operation() -> None: provider.aclose.assert_awaited_once() -async def test_any_llm_provider_uses_structured_chat_completion() -> None: +async def test_any_llm_provider_uses_json_object_with_the_strict_schema() -> None: model_result = { "headline": "Briefing", "headline_source_ids": ["source"], @@ -120,19 +130,36 @@ async def test_any_llm_provider_uses_structured_chat_completion() -> None: result = await provider.summarize("Return JSON", {"input": "数据"}) - assert client.calls == [ - { - "model": "requested-model", - "messages": [ - {"role": "system", "content": "Return JSON"}, - {"role": "user", "content": '{"input":"数据"}'}, - ], - "response_format": LLMStructuredOutput, - "temperature": 0.2, - "max_tokens": 4096, - } - ] assert result == model_result + call = client.calls[0] + assert call["response_format"] == {"type": "json_object"} + messages = call["messages"] + assert messages[0] == {"role": "system", "content": "Return JSON"} + user_content = messages[1]["content"] + assert user_content.startswith('{"input":"数据"}\n\nReturn only a JSON object') + schema = json.loads(user_content.rsplit("\n", 1)[1]) + assert schema == LLMStructuredOutput.model_json_schema() + + +async def test_json_object_transport_requires_a_final_user_message() -> None: + client = _CompletionClientStub(SimpleNamespace()) + provider = AnyLLMStructuredProvider( + client, + provider="openai", + model="requested-model", + max_output_tokens=4096, + ) + + with pytest.raises(ValueError, match="requires a final user message"): + await provider._complete( + [{"role": "system", "content": "Return JSON"}], + response_format=LLMStructuredOutput, + temperature=0.2, + max_tokens=4096, + request_error_message="LLM request failed", + ) + + assert client.calls == [] async def test_any_llm_provider_translates_service_status_with_a_narrow_schema() -> None: @@ -161,15 +188,13 @@ async def test_any_llm_provider_translates_service_status_with_a_narrow_schema() ) assert result == ("API incident", "API error rates are elevated.") - assert client.calls[0]["response_format"] is ServiceStatusTranslationOutput + assert client.calls[0]["response_format"] == {"type": "json_object"} assert client.calls[0]["temperature"] == 0.0 assert client.calls[0]["max_tokens"] == 2048 messages = client.calls[0]["messages"] - assert isinstance(messages, list) - assert messages[1] == { - "role": "user", - "content": '{"title":"API 服务异常","body":"API 服务错误率升高。"}', - } + user_content = messages[1]["content"] + assert user_content.startswith('{"title":"API 服务异常","body":"API 服务错误率升高。"}') + assert json.loads(user_content.rsplit("\n", 1)[1]) == ServiceStatusTranslationOutput.model_json_schema() async def test_any_llm_provider_assesses_notification_value_with_a_narrow_schema() -> None: @@ -197,9 +222,11 @@ async def test_any_llm_provider_assesses_notification_value_with_a_narrow_schema ) assert not result.should_notify - assert client.calls[0]["response_format"] is NotificationDecisionOutput + assert client.calls[0]["response_format"] == {"type": "json_object"} assert client.calls[0]["temperature"] == 0.0 assert client.calls[0]["max_tokens"] == 256 + user_content = client.calls[0]["messages"][1]["content"] + assert json.loads(user_content.rsplit("\n", 1)[1]) == NotificationDecisionOutput.model_json_schema() @pytest.mark.parametrize( @@ -213,7 +240,7 @@ async def test_any_llm_provider_assesses_notification_value_with_a_narrow_schema "LLM request failed", ), ( - "anthropic", + "deepseek", _anthropic_bad_request, "summarize", ("Return JSON", {"input": "data"}), @@ -227,7 +254,7 @@ async def test_any_llm_provider_assesses_notification_value_with_a_narrow_schema "LLM request failed", ), ( - "gemini", + "openrouter", _provider_status_error, "summarize", ("Return JSON", {"input": "data"}), @@ -363,7 +390,11 @@ async def test_provider_native_request_error_switches_to_fallback(monkeypatch) - assert len(fallback_client.calls) == 1 -async def test_factory_accepts_every_any_llm_completion_provider(monkeypatch) -> None: +@pytest.mark.parametrize( + "provider", + tuple(AnyLLM.get_supported_providers()), +) +def test_factory_classifies_every_any_llm_provider(monkeypatch, provider: str) -> None: created: list[tuple[str, dict[str, object]]] = [] def fake_create(provider: str, **options: object) -> _CompletionClientStub: @@ -371,22 +402,30 @@ def fake_create(provider: str, **options: object) -> _CompletionClientStub: return _CompletionClientStub(SimpleNamespace()) monkeypatch.setattr(AnyLLM, "create", fake_create) - completion_providers = [ - provider - for provider in AnyLLM.get_supported_providers() - if AnyLLM.get_provider_class(provider).SUPPORTS_COMPLETION - ] - adapters = [create_any_llm_provider(provider, "model", 1024) for provider in completion_providers] - - assert [adapter.provider for adapter in adapters] == completion_providers - assert [provider for provider, _ in created] == completion_providers - assert all( - "default_headers" not in options and "http_client" not in options and "max_retries" not in options - for _, options in created - ) - - -@pytest.mark.parametrize("provider", ("anthropic", "deepseek")) + if not AnyLLM.get_provider_class(provider).SUPPORTS_COMPLETION: + with pytest.raises( + ValueError, + match=f"any-llm provider does not support completion: {provider}", + ): + create_any_llm_provider(provider, "model", 1024) + assert created == [] + return + if provider in UNSUPPORTED_JSON_OBJECT_PROVIDERS: + with pytest.raises( + ValueError, + match=f"any-llm provider does not support required JSON Object output: {provider}", + ): + create_any_llm_provider(provider, "model", 1024) + assert created == [] + return + + adapter = create_any_llm_provider(provider, "model", 1024) + + assert adapter.provider == provider + assert created == [(provider, {"api_key": None, "api_base": None})] + + +@pytest.mark.parametrize("provider", ("openai", "deepseek")) def test_factory_passes_configured_headers_through_client_args(monkeypatch, provider: str) -> None: created: list[tuple[str, dict[str, object]]] = [] @@ -411,16 +450,18 @@ def fake_create(provider: str, **options: object) -> _CompletionClientStub: ] -def test_factory_uses_the_canonical_provider_for_header_validation(monkeypatch) -> None: +def test_factory_uses_the_canonical_provider_for_json_object_validation(monkeypatch) -> None: create = Mock() monkeypatch.setattr(AnyLLM, "create", create) - with pytest.raises(ValueError, match="Custom headers are not supported for any-llm provider: mistral"): + with pytest.raises( + ValueError, + match="any-llm provider does not support required JSON Object output: mistral", + ): create_any_llm_provider( "MISTRAL", "model", 1024, - extra_headers={"User-Agent": "weather-briefing/1"}, ) create.assert_not_called() @@ -562,11 +603,6 @@ def __init__(self) -> None: client.aclose.assert_not_awaited() -async def test_factory_rejects_provider_without_completion() -> None: - with pytest.raises(ValueError, match="does not support completion"): - create_any_llm_provider("voyage", "model", 1024) - - @pytest.mark.parametrize("provider_name", ("deepseek", "openai", "openrouter")) async def test_openai_compatible_providers_send_configured_headers( monkeypatch, @@ -650,6 +686,76 @@ def init_client( assert private_header_value not in caplog.text +async def test_openai_compatible_provider_sends_json_object_with_the_application_schema(monkeypatch) -> None: + requests: list[httpx.Request] = [] + model_result = { + "headline": "Briefing", + "headline_source_ids": ["source"], + "conclusions": [], + "active_warnings": [], + "resolved_warning_ids": [], + "advice": [], + "disaster_tracking": [], + "should_publish": True, + } + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response( + 200, + json={ + "id": "completion-id", + "object": "chat.completion", + "created": 1, + "model": "requested-model", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": { + "role": "assistant", + "content": json.dumps(model_result), + }, + } + ], + }, + ) + + def init_client( + sdk_provider: BaseOpenAIProvider, + api_key: str | None = None, + api_base: str | None = None, + **_: object, + ) -> None: + sdk_provider.client = AsyncOpenAI( + api_key=api_key, + base_url=api_base, + http_client=LoggedAsyncClient(transport=httpx.MockTransport(handler)), + ) + + monkeypatch.setattr(BaseOpenAIProvider, "_init_client", init_client) + provider = create_any_llm_provider( + "openai", + "requested-model", + 4096, + api_key="runtime-key", + api_base="https://api.example.invalid", + ) + + try: + result = await provider.summarize("Return JSON", {"input": "data"}) + finally: + await provider.aclose() + + assert result == model_result + request_body = json.loads(requests[0].content) + assert request_body["response_format"] == {"type": "json_object"} + assert "max_tokens" not in request_body + assert request_body["max_completion_tokens"] == 4096 + user_content = request_body["messages"][-1]["content"] + assert json.loads(user_content.rsplit("\n", 1)[1]) == LLMStructuredOutput.model_json_schema() + + async def test_any_llm_deepseek_uses_injected_logged_http_client(caplog) -> None: requests: list[httpx.Request] = [] model_result = { diff --git a/tests/test_cli.py b/tests/test_cli.py index c13c9ab6..0538946a 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -1200,16 +1200,16 @@ async def test_deepseek_without_base_url(self, monkeypatch) -> None: assert calls[0][0][:3] == ("deepseek", "m", 8192) assert calls[0][1]["api_base"] is None - async def test_arbitrary_any_llm_provider_is_forwarded(self, monkeypatch) -> None: + async def test_supported_any_llm_provider_is_forwarded(self, monkeypatch) -> None: calls: list[tuple[tuple[object, ...], dict[str, object]]] = [] monkeypatch.setattr( "weather_briefing.llm.any_llm.create_any_llm_provider", lambda *args, **kwargs: calls.append((args, kwargs)) or SimpleNamespace(), ) - settings = replace(_make_fake_settings(llm_provider="mistral"), api_key=None, llm_base_url=None) + settings = replace(_make_fake_settings(llm_provider="openrouter"), api_key=None, llm_base_url=None) provider = await _llm_provider(settings) assert provider is not None - assert calls[0][0][:3] == ("mistral", "m", 8192) + assert calls[0][0][:3] == ("openrouter", "m", 8192) assert calls[0][1] == { "api_key": None, "api_base": None, diff --git a/tests/test_config.py b/tests/test_config.py index 76c853df..3689e0b7 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -1043,15 +1043,15 @@ def test_positive_operational_settings_reject_zero(monkeypatch, name: str) -> No def test_any_llm_provider_uses_sdk_managed_configuration(monkeypatch) -> None: _required_environment(monkeypatch) - monkeypatch.setenv("LLM_PROVIDER", "mistral") - monkeypatch.setenv("MISTRAL_API_KEY", "generic-key") + monkeypatch.setenv("LLM_PROVIDER", "openrouter") + monkeypatch.setenv("OPENROUTER_API_KEY", "generic-key") monkeypatch.setenv("LLM_MODEL", "generic-model") - monkeypatch.setenv("MISTRAL_API_BASE", "https://compatible.example.invalid/v1") + monkeypatch.setenv("OPENROUTER_API_BASE", "https://compatible.example.invalid/v1") settings = Settings.from_env() assert settings.api_key is None - assert settings.llm_provider == "mistral" + assert settings.llm_provider == "openrouter" assert settings.llm_model == "generic-model" assert settings.llm_base_url is None @@ -1069,7 +1069,7 @@ def test_llm_fallback_provider_and_model_are_loaded(monkeypatch) -> None: def test_llm_extra_headers_are_loaded_as_immutable_mappings(monkeypatch) -> None: _required_environment(monkeypatch) - monkeypatch.setenv("LLM_PROVIDER", "anthropic") + monkeypatch.setenv("LLM_PROVIDER", "openai") monkeypatch.setenv("LLM_MODEL", "primary-model") monkeypatch.setenv("LLM_EXTRA_HEADERS", '{"User-Agent":"weather-briefing/1","X-Tenant":"primary"}') monkeypatch.setenv("LLM_FALLBACK_PROVIDER", "openrouter") @@ -1267,6 +1267,16 @@ def test_llm_provider_without_completion_raises_error(self, monkeypatch) -> None with pytest.raises(ConfigurationError, match="does not support completion"): Settings.from_env() + def test_llm_provider_without_json_object_support_raises_error(self, monkeypatch) -> None: + _required_environment(monkeypatch) + monkeypatch.setenv("LLM_PROVIDER", "mistral") + + with pytest.raises( + ConfigurationError, + match="LLM_PROVIDER does not support required JSON Object output: mistral", + ): + Settings.from_env() + @pytest.mark.parametrize( ("configured_name", "missing_name"), ( @@ -1303,6 +1313,17 @@ def test_llm_fallback_provider_without_completion_raises_error(self, monkeypatch with pytest.raises(ConfigurationError, match="LLM_FALLBACK_PROVIDER does not support completion"): Settings.from_env() + def test_llm_fallback_provider_without_json_object_support_raises_error(self, monkeypatch) -> None: + _required_environment(monkeypatch) + monkeypatch.setenv("LLM_FALLBACK_PROVIDER", "anthropic") + monkeypatch.setenv("LLM_FALLBACK_MODEL", "fallback-model") + + with pytest.raises( + ConfigurationError, + match="LLM_FALLBACK_PROVIDER does not support required JSON Object output: anthropic", + ): + Settings.from_env() + def test_invalid_float_env_value_raises_error(self, monkeypatch) -> None: _required_environment(monkeypatch) monkeypatch.setenv("HTTP_TIMEOUT_SECONDS", "not-a-number") diff --git a/tests/test_llm.py b/tests/test_llm.py index 900bb827..4879df98 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -36,7 +36,7 @@ async def acompletion( *, model: str, messages: list[dict[str, str]], - response_format: type[BaseModel], + response_format: type[BaseModel] | dict[str, object], temperature: float, max_tokens: int, ) -> object: diff --git a/weather_briefing/config/settings.py b/weather_briefing/config/settings.py index 7cd9c06d..5383b944 100644 --- a/weather_briefing/config/settings.py +++ b/weather_briefing/config/settings.py @@ -11,7 +11,10 @@ import pendulum from any_llm import AnyLLM -from ..data.any_llm_compatibility import UNSUPPORTED_DEFAULT_HEADER_PROVIDERS +from ..data.any_llm_compatibility import ( + UNSUPPORTED_DEFAULT_HEADER_PROVIDERS, + UNSUPPORTED_JSON_OBJECT_PROVIDERS, +) from ..data.bark import BARK_DEFAULT_LLM_MAX_OUTPUT_TOKENS, BARK_MAX_MESSAGE_LENGTH from ..data.service_endpoints import BARK_BASE_URL from ..models import FeedConfig, LocationSpec @@ -154,6 +157,8 @@ def from_env(cls) -> Settings: hourly_cron = cron_hour("BRIEFING_CRON", "9-23") service_status_cron = cron_expression("SERVICE_STATUS_CRON", "*/5 * * * *") llm_provider = clean_env(os.getenv("LLM_PROVIDER", "deepseek")) + llm_extra_headers = headers_from_env("LLM_EXTRA_HEADERS") + _validate_llm_headers_provider("LLM_EXTRA_HEADERS", "LLM_PROVIDER", llm_provider, llm_extra_headers) _validate_llm_provider("LLM_PROVIDER", llm_provider) llm_model = clean_env(os.getenv("LLM_MODEL")) if llm_provider == "deepseek": @@ -165,14 +170,10 @@ def from_env(cls) -> Settings: llm_base_url = None if not llm_model: raise ConfigurationError("Missing required environment variable: LLM_MODEL") - llm_extra_headers = headers_from_env("LLM_EXTRA_HEADERS") - _validate_llm_headers_provider("LLM_EXTRA_HEADERS", "LLM_PROVIDER", llm_provider, llm_extra_headers) llm_fallback_provider = clean_env(os.getenv("LLM_FALLBACK_PROVIDER")) or None llm_fallback_model = clean_env(os.getenv("LLM_FALLBACK_MODEL")) or None if (llm_fallback_provider is None) != (llm_fallback_model is None): raise ConfigurationError("LLM_FALLBACK_PROVIDER and LLM_FALLBACK_MODEL must be configured together") - if llm_fallback_provider is not None: - _validate_llm_provider("LLM_FALLBACK_PROVIDER", llm_fallback_provider) llm_fallback_extra_headers = headers_from_env("LLM_FALLBACK_EXTRA_HEADERS") if llm_fallback_extra_headers and llm_fallback_provider is None: raise ConfigurationError("LLM_FALLBACK_EXTRA_HEADERS requires LLM_FALLBACK_PROVIDER and LLM_FALLBACK_MODEL") @@ -183,6 +184,7 @@ def from_env(cls) -> Settings: llm_fallback_provider, llm_fallback_extra_headers, ) + _validate_llm_provider("LLM_FALLBACK_PROVIDER", llm_fallback_provider) locations = load_locations(locations_path) if weather_briefings_enabled else () location_ids = {location.id for location in locations} unknown_feed_locations = { @@ -306,6 +308,8 @@ def _validate_llm_provider(setting_name: str, provider: str) -> None: raise ConfigurationError(f"Unsupported {setting_name}: {provider}") if not AnyLLM.get_provider_class(provider).SUPPORTS_COMPLETION: raise ConfigurationError(f"{setting_name} does not support completion: {provider}") + if provider in UNSUPPORTED_JSON_OBJECT_PROVIDERS: + raise ConfigurationError(f"{setting_name} does not support required JSON Object output: {provider}") def _validate_llm_headers_provider( diff --git a/weather_briefing/data/any_llm_compatibility.py b/weather_briefing/data/any_llm_compatibility.py index 09ce9c3d..a4e5fbd1 100644 --- a/weather_briefing/data/any_llm_compatibility.py +++ b/weather_briefing/data/any_llm_compatibility.py @@ -1,5 +1,28 @@ """Compatibility metadata for the pinned any-llm SDK.""" +UNSUPPORTED_JSON_OBJECT_PROVIDERS = frozenset( + { + "anthropic", + "azure", + "azureanthropic", + "bedrock", + "cerebras", + "cohere", + "gemini", + "groq", + "huggingface", + "lmstudio", + "mistral", + "ollama", + "sagemaker", + "together", + "vertexai", + "vertexaianthropic", + "watsonx", + "xai", + } +) + UNSUPPORTED_DEFAULT_HEADER_PROVIDERS = frozenset( { "azure", diff --git a/weather_briefing/llm/any_llm.py b/weather_briefing/llm/any_llm.py index 0afd445f..383ae1c9 100644 --- a/weather_briefing/llm/any_llm.py +++ b/weather_briefing/llm/any_llm.py @@ -2,18 +2,23 @@ from __future__ import annotations +import json import logging from collections.abc import Iterator, Mapping from contextlib import contextmanager from inspect import isawaitable -from typing import Protocol +from typing import Any, Protocol, TypeAlias from any_llm import AnyLLM from any_llm.exceptions import AnyLLMError, LengthFinishReasonError +from any_llm.types.completion import ChatCompletionMessage from pydantic import BaseModel, ValidationError from ..api_client import api_call_context -from ..data.any_llm_compatibility import UNSUPPORTED_DEFAULT_HEADER_PROVIDERS +from ..data.any_llm_compatibility import ( + UNSUPPORTED_DEFAULT_HEADER_PROVIDERS, + UNSUPPORTED_JSON_OBJECT_PROVIDERS, +) from ..data.prompts import NOTIFICATION_POLICY from ..notifications import NotificationDecision from .base import LLMOutputLimitError, LLMRequestError, SensitiveLLMDiagnostics, serialize_llm_payload @@ -28,6 +33,8 @@ _LOGGER = logging.getLogger("weather_briefing.llm") +ResponseFormat: TypeAlias = dict[str, Any] + class LLMCompletionClient(Protocol): """Expose the any-llm completion operation used by the application adapter.""" @@ -37,7 +44,7 @@ async def acompletion( *, model: str, messages: list[dict[str, str]], - response_format: type[BaseModel], + response_format: ResponseFormat, temperature: float, max_tokens: int, ) -> object: @@ -101,6 +108,7 @@ async def _complete( max_tokens: int, request_error_message: str, ) -> object: + request_messages, request_response_format = _structured_output_request(messages, response_format) with ( api_call_context(self._provider, "chat-completions"), _normalize_request_errors( @@ -108,10 +116,22 @@ async def _complete( normalize_completion_errors=self._normalize_completion_errors, ), ): + if isinstance(self._client, AnyLLM): + any_llm_messages: list[dict[str, Any] | ChatCompletionMessage] = [ + dict(message) for message in request_messages + ] + return await self._client.acompletion( + model=self._model, + messages=any_llm_messages, + response_format=request_response_format, + stream=False, + temperature=temperature, + max_tokens=max_tokens, + ) return await self._client.acompletion( model=self._model, - messages=[*messages], - response_format=response_format, + messages=request_messages, + response_format=request_response_format, temperature=temperature, max_tokens=max_tokens, ) @@ -311,6 +331,8 @@ def create_any_llm_provider( raise ValueError(f"any-llm provider does not support completion: {canonical_provider}") if extra_headers and canonical_provider in UNSUPPORTED_DEFAULT_HEADER_PROVIDERS: raise ValueError(f"Custom headers are not supported for any-llm provider: {canonical_provider}") + if canonical_provider in UNSUPPORTED_JSON_OBJECT_PROVIDERS: + raise ValueError(f"any-llm provider does not support required JSON Object output: {canonical_provider}") client_args: dict[str, object] = {"api_key": api_key, "api_base": api_base} if extra_headers: client_args["default_headers"] = extra_headers @@ -324,3 +346,23 @@ def create_any_llm_provider( owns_client=True, normalize_completion_errors=True, ) + + +def _structured_output_request( + messages: list[dict[str, str]], + response_format: type[BaseModel], +) -> tuple[list[dict[str, str]], ResponseFormat]: + """Prepare prompt-constrained JSON Object transport.""" + if not messages or messages[-1].get("role") != "user": + raise ValueError("JSON Object structured output requires a final user message") + schema = json.dumps(response_format.model_json_schema(), ensure_ascii=False, separators=(",", ":")) + final_message = { + **messages[-1], + "content": ( + f"{messages[-1]['content']}\n\n" + "Return only a JSON object matching this JSON Schema exactly. " + "Do not wrap it in Markdown fences.\n" + f"{schema}" + ), + } + return [*messages[:-1], final_message], {"type": "json_object"} From 14c90971def6d596beafd9ab7a8043eb335a7afe Mon Sep 17 00:00:00 2001 From: IceCodeNew <32576256+IceCodeNew@users.noreply.github.com> Date: Sun, 26 Jul 2026 02:41:17 +0800 Subject: [PATCH 2/2] refactor(llm): keep SDK boundary types explicit --- tests/test_any_llm_provider.py | 21 +++++++++++++++++++++ weather_briefing/llm/any_llm.py | 15 +++++++-------- 2 files changed, 28 insertions(+), 8 deletions(-) diff --git a/tests/test_any_llm_provider.py b/tests/test_any_llm_provider.py index f14157c6..409ea9b1 100644 --- a/tests/test_any_llm_provider.py +++ b/tests/test_any_llm_provider.py @@ -162,6 +162,27 @@ async def test_json_object_transport_requires_a_final_user_message() -> None: assert client.calls == [] +async def test_json_object_transport_requires_final_user_message_content() -> None: + client = _CompletionClientStub(SimpleNamespace()) + provider = AnyLLMStructuredProvider( + client, + provider="openai", + model="requested-model", + max_output_tokens=4096, + ) + + with pytest.raises(ValueError, match="must include string content"): + await provider._complete( + [{"role": "user"}], + response_format=LLMStructuredOutput, + temperature=0.2, + max_tokens=4096, + request_error_message="LLM request failed", + ) + + assert client.calls == [] + + async def test_any_llm_provider_translates_service_status_with_a_narrow_schema() -> None: client = _CompletionClientStub( SimpleNamespace( diff --git a/weather_briefing/llm/any_llm.py b/weather_briefing/llm/any_llm.py index 383ae1c9..c36b1a10 100644 --- a/weather_briefing/llm/any_llm.py +++ b/weather_briefing/llm/any_llm.py @@ -7,11 +7,10 @@ from collections.abc import Iterator, Mapping from contextlib import contextmanager from inspect import isawaitable -from typing import Any, Protocol, TypeAlias +from typing import Protocol, TypeAlias from any_llm import AnyLLM from any_llm.exceptions import AnyLLMError, LengthFinishReasonError -from any_llm.types.completion import ChatCompletionMessage from pydantic import BaseModel, ValidationError from ..api_client import api_call_context @@ -33,7 +32,7 @@ _LOGGER = logging.getLogger("weather_briefing.llm") -ResponseFormat: TypeAlias = dict[str, Any] +ResponseFormat: TypeAlias = dict[str, object] class LLMCompletionClient(Protocol): @@ -117,12 +116,9 @@ async def _complete( ), ): if isinstance(self._client, AnyLLM): - any_llm_messages: list[dict[str, Any] | ChatCompletionMessage] = [ - dict(message) for message in request_messages - ] return await self._client.acompletion( model=self._model, - messages=any_llm_messages, + messages=[dict(message) for message in request_messages], response_format=request_response_format, stream=False, temperature=temperature, @@ -355,11 +351,14 @@ def _structured_output_request( """Prepare prompt-constrained JSON Object transport.""" if not messages or messages[-1].get("role") != "user": raise ValueError("JSON Object structured output requires a final user message") + content = messages[-1].get("content") + if not isinstance(content, str): + raise ValueError("JSON Object structured output final user message must include string content") schema = json.dumps(response_format.model_json_schema(), ensure_ascii=False, separators=(",", ":")) final_message = { **messages[-1], "content": ( - f"{messages[-1]['content']}\n\n" + f"{content}\n\n" "Return only a JSON object matching this JSON Schema exactly. " "Do not wrap it in Markdown fences.\n" f"{schema}"