diff --git a/docs/design.md b/docs/design.md index 3248c5b4..89505d0f 100644 --- a/docs/design.md +++ b/docs/design.md @@ -173,6 +173,8 @@ Open-Meteo 的逐小时空气质量和花粉预报按目标日峰值生成生活 `LLM_PROVIDER` 使用 any-llm 的 provider ID,`LLM_MODEL` 使用对应模型 ID。已部署的 DeepSeek 旧变量只在配置入口作为通用变量的后备。 +`LLM_EXTRA_HEADERS` 和 `LLM_FALLBACK_EXTRA_HEADERS` 在配置入口解析为不可变的 HTTP header 映射,再由 any-llm 适配器作为 provider client 的 `default_headers` 构造参数传入。配置入口拒绝锁定 SDK 版本中不接受这个参数的 provider;未配置 header 时不传入该参数。 + 开发环境安装 `any-llm-sdk[all]`,用于验证所有 completion provider 的装载边界。基础运行依赖只包含 SDK 核心包。官方镜像额外安装 DeepSeek、OpenAI 和 OpenRouter 所需组件。 `LLMStructuredOutput` 同时用于 SDK 的结构化输出和应用侧复验。应用还会检查来源 ID、必填建议、预警 ID 和章节间重复等领域规则。 diff --git a/docs/notes.md b/docs/notes.md index a83b6606..3b7f4265 100644 --- a/docs/notes.md +++ b/docs/notes.md @@ -48,6 +48,12 @@ 如果以后发布明确允许删除旧配置的重大版本,并且迁移说明已经给已有部署留出足够时间,就可以移除这两个后备变量。在此之前,修改模型配置边界时必须保留并测试这一优先级。 +## LLM 自定义 header 复用 provider client 参数 + +[mozilla-ai/any-llm#707](https://github.com/mozilla-ai/any-llm/pull/707) 为 any-llm 的无状态 completion API 增加了 `client_args`,并把其中的 provider-specific 参数展开传给 `AnyLLM.create()`。应用需要跨多次请求持有并显式关闭 provider client,因此不调用每次重新创建 client 的无状态 API,而是直接复用它的底层通道:把解析后的 header 映射放入本地 `client_args["default_headers"]`,再以 `AnyLLM.create(provider, **client_args)` 创建受应用管理的 provider。代码没有调用 #707 新增的具名参数,但使用的是同一条 provider client 参数透传路径。 + +升级 any-llm 或 provider SDK、新增 provider,或者 any-llm 提供统一且带能力声明的 header 接口时,必须重新检查 provider client 构造函数和实际请求,更新黑名单或改用统一接口。 + ## FallbackLLMProvider 在进程内保持粘性 这是一个有意保留的自定义外部服务集成。选择 `FallbackLLMProvider` 的原因是 any-llm 只统一调用单个 provider,不编排跨 provider 的故障切换。包装器捕获主适配器的 `LLMRequestError`,切换后在剩余生命周期内固定使用备用适配器,使同一进程里的契约修复不会回到刚刚失败的主服务。 diff --git a/env.example b/env.example index d5ea2819..98c7846d 100644 --- a/env.example +++ b/env.example @@ -6,10 +6,16 @@ # bases use the environment names documented by any-llm for that provider. LLM_PROVIDER=deepseek LLM_MODEL=deepseek-v4-flash +# Optional JSON object of request headers for the primary model. +# Values may contain credentials. Configured fields override SDK defaults, +# including Authorization and User-Agent, and are never written to logs. +# LLM_EXTRA_HEADERS={"User-Agent":"weather-briefing/1"} # Optional request-failure fallback, disabled by default. Configure both values # or neither. # LLM_FALLBACK_PROVIDER=openai # LLM_FALLBACK_MODEL=gpt-5-mini +# Optional fallback-only headers; requires both fallback values above. +# LLM_FALLBACK_EXTRA_HEADERS={"User-Agent":"weather-briefing-fallback/1"} # DeepSeek example. DEEPSEEK_API_BASE is optional. DEEPSEEK_API_KEY=replace-in-runtime-environment # DEEPSEEK_API_BASE= diff --git a/tests/test_any_llm_compatibility.py b/tests/test_any_llm_compatibility.py new file mode 100644 index 00000000..35245991 --- /dev/null +++ b/tests/test_any_llm_compatibility.py @@ -0,0 +1,27 @@ +from any_llm import AnyLLM + +from weather_briefing.data.any_llm_compatibility import UNSUPPORTED_DEFAULT_HEADER_PROVIDERS + + +def test_default_header_provider_compatibility_matches_the_pinned_sdk() -> None: + assert { + "azure", + "bedrock", + "cohere", + "gemini", + "huggingface", + "lmstudio", + "mistral", + "ollama", + "sagemaker", + "vertexai", + "watsonx", + "xai", + } == UNSUPPORTED_DEFAULT_HEADER_PROVIDERS + completion_providers = { + provider + for provider in AnyLLM.get_supported_providers() + if AnyLLM.get_provider_class(provider).SUPPORTS_COMPLETION + } + + assert completion_providers > UNSUPPORTED_DEFAULT_HEADER_PROVIDERS diff --git a/tests/test_any_llm_provider.py b/tests/test_any_llm_provider.py index 379e69c2..91808ccd 100644 --- a/tests/test_any_llm_provider.py +++ b/tests/test_any_llm_provider.py @@ -1,11 +1,14 @@ import json import logging +from collections.abc import Mapping from types import SimpleNamespace -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, Mock import httpx import pytest from any_llm import AnyLLM +from any_llm.providers.openai.base import BaseOpenAIProvider +from openai import AsyncOpenAI from pydantic import BaseModel from weather_briefing.api_client import LoggedAsyncClient @@ -190,7 +193,65 @@ def fake_create(provider: str, **options: object) -> _CompletionClientStub: assert [adapter.provider for adapter in adapters] == completion_providers assert [provider for provider, _ in created] == completion_providers - assert all("http_client" not in options and "max_retries" not in options for _, options in created) + 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")) +def test_factory_passes_configured_headers_through_client_args(monkeypatch, provider: str) -> None: + created: list[tuple[str, dict[str, object]]] = [] + + def fake_create(provider: str, **options: object) -> _CompletionClientStub: + created.append((provider, options)) + return _CompletionClientStub(SimpleNamespace()) + + monkeypatch.setattr(AnyLLM, "create", fake_create) + headers = {"User-Agent": "weather-briefing/1", "X-Tenant": "test"} + + create_any_llm_provider(provider, "model", 1024, extra_headers=headers) + + assert created == [ + ( + provider, + { + "api_key": None, + "api_base": None, + "default_headers": headers, + }, + ) + ] + + +def test_factory_uses_the_canonical_provider_for_header_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"): + create_any_llm_provider( + "MISTRAL", + "model", + 1024, + extra_headers={"User-Agent": "weather-briefing/1"}, + ) + + create.assert_not_called() + + +def test_factory_rejects_headers_for_an_unsupported_provider(monkeypatch) -> None: + create = Mock() + monkeypatch.setattr(AnyLLM, "create", create) + + with pytest.raises(ValueError, match="Custom headers are not supported for any-llm provider: mistral"): + create_any_llm_provider( + "mistral", + "model", + 1024, + extra_headers={"User-Agent": "weather-briefing/1"}, + ) + + create.assert_not_called() async def test_factory_owned_provider_closes_underlying_sdk_clients(monkeypatch) -> None: @@ -319,6 +380,89 @@ async def test_factory_rejects_provider_without_completion() -> None: 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, + caplog, + provider_name: str, +) -> None: + requests: list[httpx.Request] = [] + private_header_name = "X-Private-Token" + private_header_value = "private-value" + 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, + *, + default_headers: Mapping[str, str] | None = None, + **_: object, + ) -> None: + sdk_provider.client = AsyncOpenAI( + api_key=api_key, + base_url=api_base, + default_headers=default_headers, + http_client=LoggedAsyncClient(transport=httpx.MockTransport(handler)), + ) + + monkeypatch.setattr(BaseOpenAIProvider, "_init_client", init_client) + caplog.set_level(logging.INFO, logger="weather_briefing.api_client") + provider = create_any_llm_provider( + provider_name, + "requested-model", + 4096, + api_key="runtime-key", + api_base="https://api.example.invalid", + extra_headers={ + "User-Agent": "weather-briefing-test/1", + private_header_name: private_header_value, + }, + ) + + try: + result = await provider.summarize("Return JSON", {"input": "data"}) + finally: + await provider.aclose() + + assert result == model_result + assert requests[0].headers["user-agent"] == "weather-briefing-test/1" + assert requests[0].headers[private_header_name] == private_header_value + assert private_header_name not in caplog.text + assert private_header_value not in caplog.text + + 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 926a0cf3..c13c9ab6 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -764,8 +764,10 @@ async def fail_run( llm_provider="deepseek", llm_model="m", llm_base_url=None, + llm_extra_headers={}, llm_fallback_provider=None, llm_fallback_model=None, + llm_fallback_extra_headers={}, llm_max_output_tokens=8192, llm_max_attempts=3, http_timeout_seconds=30.0, @@ -1180,6 +1182,7 @@ async def test_deepseek_with_custom_base_url( { "api_key": "k", "api_base": "https://custom.example.invalid", + "extra_headers": {}, "diagnostics": diagnostics, }, ) @@ -1210,6 +1213,7 @@ async def test_arbitrary_any_llm_provider_is_forwarded(self, monkeypatch) -> Non assert calls[0][1] == { "api_key": None, "api_base": None, + "extra_headers": {}, "diagnostics": None, } @@ -1226,8 +1230,10 @@ def create_provider(*args, **kwargs): monkeypatch.setattr("weather_briefing.llm.any_llm.create_any_llm_provider", create_provider) settings = replace( _make_fake_settings(), + llm_extra_headers={"User-Agent": "weather-briefing/1"}, llm_fallback_provider="openai", llm_fallback_model="gpt-fallback", + llm_fallback_extra_headers={"X-Tenant": "fallback"}, ) provider = await _llm_provider(settings) @@ -1237,9 +1243,11 @@ def create_provider(*args, **kwargs): assert calls[1] == ( ("openai", "gpt-fallback", 8192), { + "extra_headers": {"X-Tenant": "fallback"}, "diagnostics": None, }, ) + assert calls[0][1]["extra_headers"] == {"User-Agent": "weather-briefing/1"} async def test_fallback_creation_failure_closes_primary_provider(self, monkeypatch) -> None: primary = AsyncMock() diff --git a/tests/test_config.py b/tests/test_config.py index 79277205..76c853df 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -61,8 +61,10 @@ def test_mainland_weather_providers_default_to_qweather_then_open_meteo(monkeypa assert [feed.id for feed in settings.feeds] == ["authority-weather"] assert settings.llm_provider == "deepseek" assert settings.llm_base_url is None + assert settings.llm_extra_headers == {} assert settings.llm_fallback_provider is None assert settings.llm_fallback_model is None + assert settings.llm_fallback_extra_headers == {} assert settings.llm_max_attempts == 3 assert settings.qweather_jwt_lifetime_seconds == 900 assert settings.llm_history_max_documents == 8 @@ -1065,6 +1067,61 @@ def test_llm_fallback_provider_and_model_are_loaded(monkeypatch) -> None: assert settings.llm_fallback_model == "gpt-fallback" +def test_llm_extra_headers_are_loaded_as_immutable_mappings(monkeypatch) -> None: + _required_environment(monkeypatch) + monkeypatch.setenv("LLM_PROVIDER", "anthropic") + 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") + monkeypatch.setenv("LLM_FALLBACK_MODEL", "fallback-model") + monkeypatch.setenv("LLM_FALLBACK_EXTRA_HEADERS", '{"User-Agent":"weather-briefing-fallback/1"}') + + settings = Settings.from_env() + + assert settings.llm_extra_headers == { + "User-Agent": "weather-briefing/1", + "X-Tenant": "primary", + } + assert settings.llm_fallback_extra_headers == { + "User-Agent": "weather-briefing-fallback/1", + } + assert not hasattr(settings.llm_extra_headers, "__setitem__") + + +def test_llm_extra_headers_reject_an_unsupported_primary_provider(monkeypatch) -> None: + _required_environment(monkeypatch) + monkeypatch.setenv("LLM_PROVIDER", "mistral") + monkeypatch.setenv("LLM_MODEL", "generic-model") + monkeypatch.setenv("LLM_EXTRA_HEADERS", '{"User-Agent":"weather-briefing/1"}') + + with pytest.raises( + ConfigurationError, + match="LLM_EXTRA_HEADERS does not support LLM_PROVIDER=mistral", + ): + Settings.from_env() + + +def test_llm_extra_headers_reject_an_unsupported_fallback_provider(monkeypatch) -> None: + _required_environment(monkeypatch) + monkeypatch.setenv("LLM_FALLBACK_PROVIDER", "mistral") + monkeypatch.setenv("LLM_FALLBACK_MODEL", "generic-model") + monkeypatch.setenv("LLM_FALLBACK_EXTRA_HEADERS", '{"User-Agent":"weather-briefing/1"}') + + with pytest.raises( + ConfigurationError, + match="LLM_FALLBACK_EXTRA_HEADERS does not support LLM_FALLBACK_PROVIDER=mistral", + ): + Settings.from_env() + + +def test_fallback_headers_require_a_configured_fallback(monkeypatch) -> None: + _required_environment(monkeypatch) + monkeypatch.setenv("LLM_FALLBACK_EXTRA_HEADERS", '{"User-Agent":"weather-briefing/1"}') + + with pytest.raises(ConfigurationError, match="requires LLM_FALLBACK_PROVIDER and LLM_FALLBACK_MODEL"): + Settings.from_env() + + def test_deepseek_model_name_remains_compatible(monkeypatch) -> None: _required_environment(monkeypatch) monkeypatch.delenv("LLM_MODEL") diff --git a/tests/test_http_headers.py b/tests/test_http_headers.py new file mode 100644 index 00000000..a6e4035e --- /dev/null +++ b/tests/test_http_headers.py @@ -0,0 +1,40 @@ +import json + +import pytest + +from weather_briefing.config.base import ConfigurationError +from weather_briefing.config.http_headers import headers_from_env + + +@pytest.mark.parametrize( + ("configured", "message"), + ( + ("not-json", "valid JSON object"), + ("null", "must be a JSON object"), + ("[]", "must be a JSON object"), + ('"value"', "valid JSON object"), + ('{"X-Count":1}', "header values must be strings"), + ('{"Bad Name":"value"}', "invalid HTTP header name"), + ('{"":"value"}', "invalid HTTP header name"), + ('{"X-Test":"café"}', "only ASCII characters"), + ('{"X-Test":"line\\nbreak"}', "control characters"), + ('{"X-Test":"value","x-test":"other"}', "duplicate HTTP header names"), + ), +) +def test_headers_from_env_rejects_invalid_json_objects(monkeypatch, configured: str, message: str) -> None: + monkeypatch.setenv("LLM_EXTRA_HEADERS", configured) + + with pytest.raises(ConfigurationError, match=message): + headers_from_env("LLM_EXTRA_HEADERS") + + +def test_header_errors_do_not_disclose_header_data(monkeypatch) -> None: + private_name = "X-Private-Token" + private_value = "private-value\nsecond-line" + monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({private_name: private_value})) + + with pytest.raises(ConfigurationError) as error: + headers_from_env("LLM_EXTRA_HEADERS") + + assert private_name not in str(error.value) + assert private_value not in str(error.value) diff --git a/weather_briefing/composition/providers.py b/weather_briefing/composition/providers.py index b8efab20..28f5c1d2 100644 --- a/weather_briefing/composition/providers.py +++ b/weather_briefing/composition/providers.py @@ -54,6 +54,7 @@ async def llm_provider( settings.llm_max_output_tokens, api_key=settings.api_key, api_base=settings.llm_base_url, + extra_headers=settings.llm_extra_headers, diagnostics=diagnostics, ) if settings.llm_fallback_provider is None or settings.llm_fallback_model is None: @@ -64,6 +65,7 @@ async def llm_provider( settings.llm_fallback_provider, settings.llm_fallback_model, settings.llm_max_output_tokens, + extra_headers=settings.llm_fallback_extra_headers, diagnostics=diagnostics, ) stack.push_async_callback(fallback.aclose) diff --git a/weather_briefing/config/http_headers.py b/weather_briefing/config/http_headers.py new file mode 100644 index 00000000..1bc487c5 --- /dev/null +++ b/weather_briefing/config/http_headers.py @@ -0,0 +1,49 @@ +"""Strict parsing for configured HTTP request headers.""" + +from __future__ import annotations + +import json +import os +import re +from collections.abc import Mapping +from types import MappingProxyType + +from .base import ConfigurationError +from .environment import clean_env + +_HEADER_NAME = re.compile(r"[!#$%&'*+\-.^_`|~0-9A-Za-z]+") + + +class _JSONObject(list[tuple[str, object]]): + """Preserve JSON object entries so duplicate names can be rejected.""" + + +def headers_from_env(name: str) -> Mapping[str, str]: + """Read an immutable JSON object containing valid HTTP header fields.""" + configured = clean_env(os.getenv(name)) + if not configured: + return MappingProxyType({}) + try: + payload = json.loads(configured, object_pairs_hook=_JSONObject) + except json.JSONDecodeError as exc: + raise ConfigurationError(f"{name} must be a valid JSON object") from exc + if not isinstance(payload, _JSONObject): + raise ConfigurationError(f"{name} must be a JSON object") + + headers: dict[str, str] = {} + normalized_names: set[str] = set() + for header_name, header_value in payload: + if _HEADER_NAME.fullmatch(header_name) is None: + raise ConfigurationError(f"{name} contains an invalid HTTP header name") + normalized_name = header_name.casefold() + if normalized_name in normalized_names: + raise ConfigurationError(f"{name} contains duplicate HTTP header names") + normalized_names.add(normalized_name) + if not isinstance(header_value, str): + raise ConfigurationError(f"{name} header values must be strings") + if not header_value.isascii(): + raise ConfigurationError(f"{name} header values must contain only ASCII characters") + if any(ord(character) < 32 or ord(character) == 127 for character in header_value): + raise ConfigurationError(f"{name} header values must not contain control characters") + headers[header_name] = header_value + return MappingProxyType(headers) diff --git a/weather_briefing/config/settings.py b/weather_briefing/config/settings.py index 82bdc7bc..7cd9c06d 100644 --- a/weather_briefing/config/settings.py +++ b/weather_briefing/config/settings.py @@ -3,6 +3,7 @@ from __future__ import annotations import os +from collections.abc import Mapping from dataclasses import dataclass from pathlib import Path from urllib.parse import urlsplit @@ -10,6 +11,7 @@ import pendulum from any_llm import AnyLLM +from ..data.any_llm_compatibility import UNSUPPORTED_DEFAULT_HEADER_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 @@ -35,6 +37,7 @@ state_path_from_env, ) from .feeds import load_feeds +from .http_headers import headers_from_env from .locations import load_locations _DEFAULT_LLM_MAX_OUTPUT_TOKENS = 8192 @@ -48,8 +51,10 @@ class Settings: llm_provider: str llm_model: str llm_base_url: str | None + llm_extra_headers: Mapping[str, str] llm_fallback_provider: str | None llm_fallback_model: str | None + llm_fallback_extra_headers: Mapping[str, str] llm_max_output_tokens: int llm_max_attempts: int http_timeout_seconds: float @@ -160,12 +165,24 @@ 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") + if llm_fallback_provider is not None: + _validate_llm_headers_provider( + "LLM_FALLBACK_EXTRA_HEADERS", + "LLM_FALLBACK_PROVIDER", + llm_fallback_provider, + llm_fallback_extra_headers, + ) locations = load_locations(locations_path) if weather_briefings_enabled else () location_ids = {location.id for location in locations} unknown_feed_locations = { @@ -229,8 +246,10 @@ def from_env(cls) -> Settings: llm_provider=llm_provider, llm_model=llm_model, llm_base_url=llm_base_url.rstrip("/") if llm_base_url else None, + llm_extra_headers=llm_extra_headers, llm_fallback_provider=llm_fallback_provider, llm_fallback_model=llm_fallback_model, + llm_fallback_extra_headers=llm_fallback_extra_headers, llm_max_output_tokens=llm_max_output_tokens, llm_max_attempts=positive_integer("LLM_MAX_ATTEMPTS", 3), http_timeout_seconds=positive_float("HTTP_TIMEOUT_SECONDS", 30), @@ -287,3 +306,14 @@ 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}") + + +def _validate_llm_headers_provider( + headers_setting_name: str, + provider_setting_name: str, + provider: str, + headers: Mapping[str, str], +) -> None: + """Reject custom headers for providers without the shared client option.""" + if headers and provider in UNSUPPORTED_DEFAULT_HEADER_PROVIDERS: + raise ConfigurationError(f"{headers_setting_name} does not support {provider_setting_name}={provider}") diff --git a/weather_briefing/data/any_llm_compatibility.py b/weather_briefing/data/any_llm_compatibility.py new file mode 100644 index 00000000..09ce9c3d --- /dev/null +++ b/weather_briefing/data/any_llm_compatibility.py @@ -0,0 +1,18 @@ +"""Compatibility metadata for the pinned any-llm SDK.""" + +UNSUPPORTED_DEFAULT_HEADER_PROVIDERS = frozenset( + { + "azure", + "bedrock", + "cohere", + "gemini", + "huggingface", + "lmstudio", + "mistral", + "ollama", + "sagemaker", + "vertexai", + "watsonx", + "xai", + } +) diff --git a/weather_briefing/llm/any_llm.py b/weather_briefing/llm/any_llm.py index 0422cb1b..e00cd40a 100644 --- a/weather_briefing/llm/any_llm.py +++ b/weather_briefing/llm/any_llm.py @@ -3,6 +3,7 @@ from __future__ import annotations import logging +from collections.abc import Mapping from inspect import isawaitable from typing import Protocol @@ -11,6 +12,7 @@ from pydantic import BaseModel from ..api_client import api_call_context +from ..data.any_llm_compatibility import UNSUPPORTED_DEFAULT_HEADER_PROVIDERS from ..data.prompts import NOTIFICATION_POLICY from ..notifications import NotificationDecision from .base import LLMOutputLimitError, LLMRequestError, SensitiveLLMDiagnostics, serialize_llm_payload @@ -259,16 +261,23 @@ def create_any_llm_provider( *, api_key: str | None = None, api_base: str | None = None, + extra_headers: Mapping[str, str] | None = None, diagnostics: SensitiveLLMDiagnostics | None = None, ) -> AnyLLMStructuredProvider: """Create an application adapter for any supported any-llm completion provider.""" provider_class = AnyLLM.get_provider_class(provider) + canonical_provider = provider_class.PROVIDER_NAME if not provider_class.SUPPORTS_COMPLETION: - raise ValueError(f"any-llm provider does not support completion: {provider}") - sdk_client = AnyLLM.create(provider, api_key=api_key, api_base=api_base) + 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}") + client_args: dict[str, object] = {"api_key": api_key, "api_base": api_base} + if extra_headers: + client_args["default_headers"] = extra_headers + sdk_client = AnyLLM.create(canonical_provider, **client_args) return AnyLLMStructuredProvider( sdk_client, - provider=provider, + provider=canonical_provider, model=model, max_output_tokens=max_output_tokens, diagnostics=diagnostics,