From 572299a45eec34446559b711a589b85cf7f8e2e2 Mon Sep 17 00:00:00 2001 From: freeman Date: Thu, 16 Jul 2026 19:27:06 +0800 Subject: [PATCH] =?UTF-8?q?feat(llm):=20=E6=96=B0=E5=A2=9E=20Amazon=20Bedr?= =?UTF-8?q?ock=20=E6=8F=90=E4=BE=9B=E5=95=86=EF=BC=8C=E6=94=AF=E6=8C=81=20?= =?UTF-8?q?AK/SK=20=E4=B8=8E=20Bedrock=20API=20Key=20=E5=8F=8C=E8=AE=A4?= =?UTF-8?q?=E8=AF=81=20(#6130)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/agent/llm/helper.py | 29 +- app/agent/llm/provider.py | 489 ++++++++++++++++++++++++++++- requirements.in | 2 + tests/test_llm_provider_bedrock.py | 392 +++++++++++++++++++++++ 4 files changed, 910 insertions(+), 2 deletions(-) create mode 100644 tests/test_llm_provider_bedrock.py diff --git a/app/agent/llm/helper.py b/app/agent/llm/helper.py index 292708deb..04accebc4 100644 --- a/app/agent/llm/helper.py +++ b/app/agent/llm/helper.py @@ -1058,6 +1058,29 @@ class LLMHelper: http_async_client=_build_httpx_client(llm_proxy, async_client=True), **thinking_kwargs, ) + elif runtime["runtime"] == "bedrock": + from langchain_aws import ChatBedrockConverse + + from app.agent.llm.provider import LLMProviderManager + + aws_region = runtime.get("aws_region") or "us-east-1" + aws_auth = runtime.get("aws_auth") or {} + # Bearer 认证需要跳过 SigV4 签名并注入 Authorization 头,SigV4 认证 + # 直接以 AK/SK 签名;两种方式统一由 provider 管理器构造 boto3 客户端。 + bedrock_client = LLMProviderManager().create_bedrock_client( + "bedrock-runtime", + region=aws_region, + credentials=aws_auth, + base_url=runtime.get("base_url"), + use_proxy=use_proxy, + read_timeout=settings.LLM_TOOL_TIMEOUT, + ) + model = ChatBedrockConverse( + model_id=model_name, + client=bedrock_client, + temperature=temperature_value, + disable_streaming=not streaming, + ) elif runtime["runtime"] in {"anthropic_compatible", "copilot_anthropic"}: from langchain_anthropic import ChatAnthropic @@ -1107,7 +1130,11 @@ class LLMHelper: # 优先使用 provider / models.dev 目录中的上下文上限,减少用户手填成本。 model_profile = getattr(model, "profile", None) if model_profile: - logger.debug(f"使用LLM模型: {model.model},Profile: {model.profile}") + # ChatBedrockConverse 等模型类没有 model 属性,模型名存放在 model_id。 + logged_model_name = getattr(model, "model", None) or getattr( + model, "model_id", model_name + ) + logger.debug(f"使用LLM模型: {logged_model_name},Profile: {model_profile}") else: model_record = runtime.get("model_record") or {} model_metadata = runtime.get("model_metadata") or {} diff --git a/app/agent/llm/provider.py b/app/agent/llm/provider.py index 70bf98b65..d6afcd7fe 100644 --- a/app/agent/llm/provider.py +++ b/app/agent/llm/provider.py @@ -7,13 +7,14 @@ import base64 import copy import hashlib import json +import re import secrets import threading import time from dataclasses import dataclass, field from pathlib import Path from typing import Any, Dict, Optional, Tuple -from urllib.parse import urlencode +from urllib.parse import urlencode, urlsplit import aiofiles import httpx @@ -106,6 +107,90 @@ class LLMProviderManager(metaclass=Singleton): _MODELS_DEV_BUNDLED_PATH = Path(__file__).with_name("models.json") _MODELS_DEV_CACHE_TTL = 7 * 24 * 60 * 60 _AUTH_SESSION_DONE_RETENTION = 300 + _BEDROCK_DEFAULT_REGION = "us-east-1" + _BEDROCK_API_KEY_PREFIX = "bedrock-api-key-" + _BEDROCK_GPT_OSS_BASE_REGIONS = ( + "ap-northeast-1", + "ap-south-1", + "ap-southeast-2", + "eu-central-1", + "eu-north-1", + "eu-west-1", + "eu-west-2", + "sa-east-1", + "us-east-1", + "us-east-2", + "us-west-2", + ) + _BEDROCK_GPT_OSS_SAFEGUARD_REGIONS = ( + "ap-northeast-1", + "ap-south-1", + "ap-southeast-2", + "eu-west-1", + "eu-west-2", + "sa-east-1", + "us-east-1", + "us-east-2", + "us-west-2", + ) + _BEDROCK_ON_DEMAND_MODEL_REGIONS = { + "openai.gpt-oss-120b-1:0": _BEDROCK_GPT_OSS_BASE_REGIONS, + "openai.gpt-oss-20b-1:0": _BEDROCK_GPT_OSS_BASE_REGIONS, + "openai.gpt-oss-safeguard-120b": _BEDROCK_GPT_OSS_SAFEGUARD_REGIONS, + "openai.gpt-oss-safeguard-20b": _BEDROCK_GPT_OSS_SAFEGUARD_REGIONS, + "amazon.nova-lite-v1:0": ( + "ap-northeast-1", + "ap-southeast-2", + "eu-west-2", + "us-east-1", + "us-gov-west-1", + ), + "amazon.nova-micro-v1:0": ( + "ap-southeast-2", + "eu-west-2", + "us-east-1", + "us-gov-west-1", + ), + "amazon.nova-pro-v1:0": ( + "ap-southeast-2", + "eu-west-2", + "us-east-1", + "us-gov-west-1", + ), + "anthropic.claude-3-5-haiku-20241022-v1:0": ( + "us-west-2", + ), + "anthropic.claude-3-5-sonnet-20240620-v1:0": ( + "ap-northeast-1", + "ap-northeast-2", + "ap-southeast-1", + "eu-central-1", + "eu-central-2", + "us-east-1", + "us-gov-west-1", + "us-west-2", + ), + "anthropic.claude-3-5-sonnet-20241022-v2:0": ( + "ap-southeast-2", + "us-west-2", + ), + "anthropic.claude-3-7-sonnet-20250219-v1:0": ( + "eu-west-2", + "us-gov-west-1", + ), + "anthropic.claude-3-haiku-20240307-v1:0": ( + "ap-northeast-1", + "ap-northeast-2", + "ap-south-1", + "ap-southeast-2", + "eu-central-1", + "eu-west-1", + "eu-west-3", + "us-east-1", + "us-gov-west-1", + "us-west-2", + ), + } _CHATGPT_CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann" _CHATGPT_ISSUER = "https://auth.openai.com" _CHATGPT_CODEX_BASE_URL = "https://chatgpt.com/backend-api/codex" @@ -367,6 +452,50 @@ class LLMProviderManager(metaclass=Singleton): api_key_hint="填写 Anthropic API Key。", description="Anthropic Claude 官方端点。", ), + ProviderSpec( + id="amazon-bedrock", + name="Amazon Bedrock", + runtime="bedrock", + models_dev_provider_id="amazon-bedrock", + default_base_url="https://bedrock-runtime.us-east-1.amazonaws.com", + base_url_presets=( + url_preset( + id="bedrock-us-east-1", + label="美东(弗吉尼亚北部)us-east-1", + value="https://bedrock-runtime.us-east-1.amazonaws.com", + ), + url_preset( + id="bedrock-us-west-2", + label="美西(俄勒冈)us-west-2", + value="https://bedrock-runtime.us-west-2.amazonaws.com", + ), + url_preset( + id="bedrock-eu-central-1", + label="欧洲(法兰克福)eu-central-1", + value="https://bedrock-runtime.eu-central-1.amazonaws.com", + ), + url_preset( + id="bedrock-ap-northeast-1", + label="亚太(东京)ap-northeast-1", + value="https://bedrock-runtime.ap-northeast-1.amazonaws.com", + ), + url_preset( + id="bedrock-ap-southeast-1", + label="亚太(新加坡)ap-southeast-1", + value="https://bedrock-runtime.ap-southeast-1.amazonaws.com", + ), + ), + base_url_editable=True, + api_key_label="Bedrock API Key / AK:SK", + api_key_hint=( + "支持两种认证方式:填写 Amazon Bedrock API Key(bedrock-api-key- 开头," + "Bearer 认证);或填写 Access Key ID:Secret Access Key(可选追加 :Session Token," + "SigV4 认证)。Base URL 决定 AWS Region。" + ), + model_list_strategy="bedrock", + description="Amazon Bedrock 托管模型服务,支持 Bedrock API Key 与 AK/SK 双认证。", + sort_order=35, + ), ProviderSpec( id="deepseek", name="DeepSeek", @@ -1743,6 +1872,112 @@ class LLMProviderManager(metaclass=Singleton): return normalized[:-3] return normalized + @classmethod + def _extract_bedrock_region(cls, base_url: Optional[str]) -> str: + """ + 从 Bedrock 运行时端点 URL 中提取 AWS Region + + 兼容标准端点、FIPS 端点与 PrivateLink(VPCE)端点等主机名形态, + 从中识别 Region 段。 + + :param base_url: 形如 https://bedrock-runtime.us-east-1.amazonaws.com 的端点地址 + :return: 提取到的 Region,无法识别时回退 us-east-1 + """ + hostname = urlsplit((base_url or "").strip().lower()).hostname or "" + match = re.search( + r"(?:^|\.)(?:bedrock(?:-runtime)?(?:-fips)?)" + r"\.([a-z0-9-]+-\d+)(?:\.|$)", + hostname, + ) + if match: + return match.group(1) + return cls._BEDROCK_DEFAULT_REGION + + # Inference Profile 的地理前缀与可用 Region 的对应关系,用于降级目录按 + # 当前 Region 过滤掉不可调用的 Profile 条目。 + _BEDROCK_GEO_PREFIXES: dict[str, tuple[str, ...]] = { + "us": ("us-east-", "us-west-"), + "eu": ("eu-",), + "apac": ("ap-",), + "au": ("ap-southeast-2", "ap-southeast-4"), + "jp": ("ap-northeast-1", "ap-northeast-3"), + "ca": ("ca-",), + } + _BEDROCK_NON_COMMERCIAL_REGION_PREFIXES = ( + "cn-", + "eu-isoe-", + "us-gov-", + "us-iso-", + "us-isob-", + "us-isof-", + ) + + @classmethod + def _bedrock_model_matches_region(cls, model_id: str, region: str) -> bool: + """ + 判断目录中的模型 ID 在指定 Region 是否可调用 + + models.dev 目录同时收录裸模型 ID(直连调用)与带地理前缀的 + Inference Profile ID(us./eu./apac./global. 等)。带前缀的条目只在 + 对应地理分区和 AWS 分区的 Region 可用;global Profile 仅允许商业 + AWS 分区。裸 ID 仅在明确记录的 ON_DEMAND Region 可用,未知条目 + 按不可直连处理。 + + :param model_id: 目录中的模型 ID + :param region: 当前 Base URL 对应的 AWS Region + :return: 该模型在当前 Region 可调用时返回 True + """ + prefix = model_id.split(".", 1)[0] + if prefix == "global": + return not region.startswith(cls._BEDROCK_NON_COMMERCIAL_REGION_PREFIXES) + region_prefixes = cls._BEDROCK_GEO_PREFIXES.get(prefix) + if region_prefixes is not None: + return ( + not region.startswith(cls._BEDROCK_NON_COMMERCIAL_REGION_PREFIXES) + and region.startswith(region_prefixes) + ) + on_demand_regions = cls._BEDROCK_ON_DEMAND_MODEL_REGIONS.get(model_id) + return on_demand_regions is not None and region in on_demand_regions + + @classmethod + def _parse_bedrock_credentials(cls, api_key: Optional[str]) -> dict[str, Any]: + """ + 解析 Bedrock 凭证字符串,识别 Bearer 与 SigV4 两种认证方式 + + - Bedrock API Key(bedrock-api-key- 开头的长期 Key,或控制台生成的短期 + Token)走 Bearer 认证; + - `AccessKeyId:SecretAccessKey` 或 `AccessKeyId:SecretAccessKey:SessionToken` + 走 SigV4 认证,AWS Access Key ID 均以 "AKIA"/"ASIA" 开头。 + + :param api_key: 用户在 API Key 输入框填写的凭证内容 + :return: 含 auth_scheme 及对应凭证字段的字典 + """ + normalized = str(api_key or "").strip() + if not normalized: + raise LLMProviderAuthError( + "Amazon Bedrock 需要填写 Bedrock API Key 或 Access Key ID:Secret Access Key" + ) + + if not normalized.startswith(cls._BEDROCK_API_KEY_PREFIX): + parts = [part.strip() for part in normalized.split(":")] + if len(parts) in {2, 3} and all(parts): + credentials = { + "auth_scheme": "sigv4", + "access_key_id": parts[0], + "secret_access_key": parts[1], + } + if len(parts) == 3: + credentials["session_token"] = parts[2] + return credentials + if ":" in normalized: + raise LLMProviderAuthError( + "Amazon Bedrock AK/SK 凭证格式不正确," + "请按 AccessKeyId:SecretAccessKey 或 " + "AccessKeyId:SecretAccessKey:SessionToken 填写" + ) + + return {"auth_scheme": "bearer", "bearer_token": normalized} + async def _list_models_from_google( self, api_key: str, @@ -1857,6 +2092,235 @@ class LLMProviderManager(metaclass=Singleton): ) return sorted(results, key=lambda item: item["name"].lower()) + def _build_bedrock_boto3_config( + self, + use_proxy: Optional[bool] = None, + ) -> Any: + """ + 构造 Bedrock boto3 客户端配置,统一超时、重试与代理策略 + + :param use_proxy: 是否使用系统代理,None 时读取 LLM_USE_PROXY 配置 + :return: botocore Config 实例 + """ + from botocore.config import Config + + should_use_proxy = settings.LLM_USE_PROXY if use_proxy is None else use_proxy + proxies = None + if should_use_proxy and settings.PROXY_HOST: + proxies = {"http": settings.PROXY_HOST, "https": settings.PROXY_HOST} + return Config( + connect_timeout=10, + read_timeout=60, + retries={"max_attempts": 3, "mode": "standard"}, + proxies=proxies, + ) + + @staticmethod + def _bedrock_endpoint_url( + service_name: str, base_url: Optional[str] + ) -> Optional[str]: + """ + 解析应传给 boto3 客户端的自定义端点 URL + + 标准公有端点交由 boto3 按 Region 自行推导;用户填写 PrivateLink、 + FIPS 等非标准端点时才显式透传,保证所选网络路径实际生效。 + + :param service_name: boto3 服务名(bedrock 或 bedrock-runtime) + :param base_url: 用户配置的 Base URL + :return: 需要显式指定端点时返回 URL,否则返回 None + """ + normalized = (base_url or "").strip().rstrip("/") + if not normalized: + return None + if re.fullmatch( + rf"https://{service_name}\.[a-z0-9-]+\.amazonaws\.com", + normalized, + ): + return None + return normalized + + def create_bedrock_client( + self, + service_name: str, + region: str, + credentials: dict[str, Any], + base_url: Optional[str] = None, + use_proxy: Optional[bool] = None, + read_timeout: Optional[int] = None, + ) -> Any: + """ + 按解析后的凭证创建 Bedrock boto3 客户端,Bearer 方式注入 Authorization 头 + + :param service_name: boto3 服务名(bedrock 或 bedrock-runtime) + :param region: AWS Region + :param credentials: `_parse_bedrock_credentials` 的解析结果 + :param base_url: 用户配置的 Base URL,非标准端点(PrivateLink/FIPS 等)时透传给 boto3 + :param use_proxy: 是否使用系统代理 + :param read_timeout: 读取超时秒数,None 时使用默认值 + :return: boto3 客户端实例 + """ + import boto3 + from botocore import UNSIGNED + + config = self._build_bedrock_boto3_config(use_proxy) + if read_timeout: + config = config.merge(type(config)(read_timeout=read_timeout)) + endpoint_kwargs: dict[str, Any] = {} + endpoint_url = self._bedrock_endpoint_url(service_name, base_url) + if endpoint_url: + endpoint_kwargs["endpoint_url"] = endpoint_url + + if credentials["auth_scheme"] == "sigv4": + return boto3.client( + service_name, + region_name=region, + aws_access_key_id=credentials["access_key_id"], + aws_secret_access_key=credentials["secret_access_key"], + aws_session_token=credentials.get("session_token"), + config=config, + **endpoint_kwargs, + ) + + # Bearer 认证:以 UNSIGNED 跳过 SigV4 签名,再把 API Key 注入 Authorization 头。 + bearer_token = credentials["bearer_token"] + config = config.merge(type(config)(signature_version=UNSIGNED)) + client = boto3.client( + service_name, + region_name=region, + aws_access_key_id="unsigned", + aws_secret_access_key="unsigned", + config=config, + **endpoint_kwargs, + ) + + def _inject_bearer(request: Any, **_kwargs: Any) -> None: + request.headers["Authorization"] = f"Bearer {bearer_token}" + + client.meta.events.register( + f"request-created.{service_name}", + _inject_bearer, + ) + return client + + async def _list_models_from_bedrock_fallback( + self, + region: str, + use_proxy: Optional[bool] = None, + ) -> list[dict[str, Any]]: + """ + 从 models.dev 目录筛选当前 Region 可调用的 Bedrock 模型 + + :param region: 当前 Base URL 对应的 AWS Region + :param use_proxy: 是否使用系统代理 + :return: 过滤后的标准化模型记录列表 + """ + models = await self._list_models_from_models_dev_only( + provider_id="amazon-bedrock", + use_proxy=use_proxy, + ) + return [ + model + for model in models + if self._bedrock_model_matches_region(model["id"], region) + ] + + async def _list_models_from_bedrock( + self, + api_key: str, + base_url: Optional[str], + use_proxy: Optional[bool] = None, + ) -> list[dict[str, Any]]: + """ + 从 Bedrock 控制面拉取模型目录,聚合跨区 Inference Profile 与直连模型 + + Bedrock 多数新模型仅允许通过 Inference Profile(us./eu./apac./global. 前缀) + 调用,因此优先列出 Profile,再补充支持 ON_DEMAND 直连的基础模型。 + + :param api_key: 用户填写的凭证内容(Bedrock API Key 或 AK/SK) + :param base_url: Bedrock 运行时端点,决定 Region + :param use_proxy: 是否使用系统代理 + :return: 标准化后的模型记录列表 + """ + credentials = self._parse_bedrock_credentials(api_key) + region = self._extract_bedrock_region(base_url) + # runtime VPCE 无法安全推导对应的控制面 VPCE;FIPS 端点也不能绕回 + # 公有非 FIPS 控制面,因此直接使用本地目录。 + if self._bedrock_endpoint_url("bedrock-runtime", base_url): + return await self._list_models_from_bedrock_fallback(region, use_proxy) + client = self.create_bedrock_client( + "bedrock", + region=region, + credentials=credentials, + use_proxy=use_proxy, + ) + + def _fetch() -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: + profiles: list[dict[str, Any]] = [] + paginator = client.get_paginator("list_inference_profiles") + for page in paginator.paginate(typeEquals="SYSTEM_DEFINED"): + profiles.extend(page.get("inferenceProfileSummaries") or []) + foundation = client.list_foundation_models( + byOutputModality="TEXT", + byInferenceType="ON_DEMAND", + ).get("modelSummaries") or [] + return profiles, foundation + + try: + profile_summaries, foundation_summaries = await asyncio.to_thread(_fetch) + except Exception as err: + # 部分 Bedrock API Key 的授权范围仅覆盖 bedrock-runtime 推理接口, + # 控制面查询被拒时降级到 models.dev 目录,保证仍能选择模型。 + logger.warning( + f"获取 Amazon Bedrock 控制面模型列表失败,降级 models.dev 目录: {err}" + ) + return await self._list_models_from_bedrock_fallback(region, use_proxy) + finally: + await asyncio.to_thread(client.close) + + results: list[dict[str, Any]] = [] + seen_ids: set[str] = set() + + def _append_record(model_id: str, display_name: Optional[str]) -> None: + if not model_id or model_id in seen_ids: + return + seen_ids.add(model_id) + # Inference Profile 带区域前缀,models.dev 目录按基础模型 ID 收录, + # 去掉首个前缀段再查一次元数据。 + metadata = self._cached_models_dev_model("amazon-bedrock", model_id) + if not metadata and "." in model_id: + metadata = self._cached_models_dev_model( + "amazon-bedrock", + model_id.split(".", 1)[1], + ) + results.append( + self._normalize_model_record( + model_id=model_id, + display_name=display_name or (metadata or {}).get("name") or model_id, + metadata=metadata or {}, + source="provider", + ) + ) + + for profile in profile_summaries: + if (profile.get("status") or "ACTIVE") != "ACTIVE": + continue + _append_record( + str(profile.get("inferenceProfileId") or "").strip(), + profile.get("inferenceProfileName"), + ) + # 控制面已按当前 Region 和 ON_DEMAND 筛选,不能复用仅面向 + # models.dev 降级目录的静态白名单,否则 AWS 新增模型会被遗漏。 + for summary in foundation_summaries: + lifecycle = (summary.get("modelLifecycle") or {}).get("status") or "ACTIVE" + if lifecycle != "ACTIVE": + continue + _append_record( + str(summary.get("modelId") or "").strip(), + summary.get("modelName"), + ) + + return sorted(results, key=lambda item: item["name"].lower()) + @staticmethod def _copilot_headers( token: Optional[str] = None, include_auth: bool = True @@ -2064,6 +2528,13 @@ class LLMProviderManager(metaclass=Singleton): use_proxy=use_proxy, ) + if resolved_model_list_strategy == "bedrock": + return await self._list_models_from_bedrock( + api_key=runtime["api_key"], + base_url=runtime.get("base_url"), + use_proxy=use_proxy, + ) + if resolved_model_list_strategy == "anthropic_compatible": return await self._list_models_from_models_dev_only( provider_id=provider_id, @@ -2731,6 +3202,22 @@ class LLMProviderManager(metaclass=Singleton): ) return result + if resolved_runtime == "bedrock": + effective_base_url = normalized_base_url or self._default_base_url_for_provider( + spec + ) + credentials = self._parse_bedrock_credentials(normalized_api_key) + result.update( + { + "api_key": normalized_api_key, + "base_url": effective_base_url, + "aws_region": self._extract_bedrock_region(effective_base_url), + "aws_auth": credentials, + "auth_mode": "api_key", + } + ) + return result + if resolved_runtime == "anthropic_compatible": effective_base_url = normalized_base_url or self._default_base_url_for_provider( spec diff --git a/requirements.in b/requirements.in index ea2c1fa87..3f598b722 100644 --- a/requirements.in +++ b/requirements.in @@ -80,6 +80,8 @@ langchain~=1.3.9 langchain-core~=1.4.7 langchain-community~=0.4.2 langchain-anthropic~=1.4.6 +langchain-aws~=1.6.2 +boto3~=1.42.42 langchain-openai~=1.3.2 langchain-google-genai~=4.2.5 langchain-deepseek~=1.1.0 diff --git a/tests/test_llm_provider_bedrock.py b/tests/test_llm_provider_bedrock.py new file mode 100644 index 000000000..c3a6f2657 --- /dev/null +++ b/tests/test_llm_provider_bedrock.py @@ -0,0 +1,392 @@ +"""Amazon Bedrock provider 的凭证解析、Region 提取与运行时解析测试""" + +import asyncio +import time +from unittest.mock import MagicMock, patch + +import pytest + +from app.agent.llm.provider import ( + LLMProviderAuthError, + LLMProviderManager, +) + + +@pytest.fixture(autouse=True) +def _reset_manager_singleton(): + """每个用例前后清理 LLMProviderManager 单例,避免缓存互相污染""" + LLMProviderManager._instances.clear() + yield + LLMProviderManager._instances.clear() + + +def test_bedrock_provider_registered(): + manager = LLMProviderManager() + spec = manager.get_provider("amazon-bedrock") + + assert spec.runtime == "bedrock" + assert spec.model_list_strategy == "bedrock" + assert spec.base_url_editable is True + assert spec.default_base_url == "https://bedrock-runtime.us-east-1.amazonaws.com" + preset_ids = {preset.id for preset in spec.base_url_presets} + assert "bedrock-us-east-1" in preset_ids + assert "bedrock-ap-northeast-1" in preset_ids + + +def test_parse_bedrock_credentials_bearer_api_key(): + credentials = LLMProviderManager._parse_bedrock_credentials( + "bedrock-api-key-abcdef123456" + ) + + assert credentials["auth_scheme"] == "bearer" + assert credentials["bearer_token"] == "bedrock-api-key-abcdef123456" + + +def test_parse_bedrock_credentials_sigv4_ak_sk(): + credentials = LLMProviderManager._parse_bedrock_credentials( + "AKIAIOSFODNN7EXAMPLE:wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY" + ) + + assert credentials["auth_scheme"] == "sigv4" + assert credentials["access_key_id"] == "AKIAIOSFODNN7EXAMPLE" + assert credentials["secret_access_key"] == "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY" + assert "session_token" not in credentials + + +def test_parse_bedrock_credentials_sigv4_with_session_token(): + credentials = LLMProviderManager._parse_bedrock_credentials( + "ASIAIOSFODNN7EXAMPLE:secret/key:session-token-value" + ) + + assert credentials["auth_scheme"] == "sigv4" + assert credentials["session_token"] == "session-token-value" + + +def test_parse_bedrock_credentials_empty_rejected(): + with pytest.raises(LLMProviderAuthError): + LLMProviderManager._parse_bedrock_credentials("") + + +def test_parse_bedrock_credentials_malformed_colon_rejected(): + with pytest.raises(LLMProviderAuthError): + LLMProviderManager._parse_bedrock_credentials("AKIA123:") + + +def test_extract_bedrock_region_from_base_url(): + """应从标准、FIPS 与 PrivateLink Bedrock 端点提取 Region""" + extract = LLMProviderManager._extract_bedrock_region + + assert extract("https://bedrock-runtime.us-east-1.amazonaws.com") == "us-east-1" + assert extract("https://bedrock-runtime.ap-northeast-1.amazonaws.com/") == "ap-northeast-1" + assert extract("https://bedrock-runtime.mx-central-1.amazonaws.com") == "mx-central-1" + assert extract("https://bedrock.eu-central-1.amazonaws.com") == "eu-central-1" + # FIPS 与 PrivateLink(VPCE)端点同样能识别 Region + assert extract("https://bedrock-runtime-fips.us-east-1.amazonaws.com") == "us-east-1" + assert ( + extract("https://vpce-0abc123-xyz.bedrock-runtime.us-west-2.vpce.amazonaws.com") + == "us-west-2" + ) + # 无法识别时回退默认 Region + assert extract("https://example.com/us-west-2") == "us-east-1" + assert extract("https://example.com?region=.us-west-2.") == "us-east-1" + assert extract("https://example.com") == "us-east-1" + assert extract(None) == "us-east-1" + assert extract("") == "us-east-1" + + +def test_bedrock_endpoint_url_passthrough(): + """自定义 Bedrock 端点应透传,标准端点交由 boto3 推导""" + resolve = LLMProviderManager._bedrock_endpoint_url + + # 标准公有端点交由 boto3 推导,不显式透传 + assert resolve("bedrock-runtime", "https://bedrock-runtime.us-east-1.amazonaws.com") is None + assert resolve("bedrock", "https://bedrock.eu-central-1.amazonaws.com") is None + assert resolve("bedrock-runtime", None) is None + assert resolve("bedrock-runtime", "") is None + # FIPS / PrivateLink 等非标准端点需要显式生效 + assert ( + resolve("bedrock-runtime", "https://bedrock-runtime-fips.us-east-1.amazonaws.com") + == "https://bedrock-runtime-fips.us-east-1.amazonaws.com" + ) + assert ( + resolve( + "bedrock-runtime", + "https://vpce-0abc123-xyz.bedrock-runtime.us-west-2.vpce.amazonaws.com/", + ) + == "https://vpce-0abc123-xyz.bedrock-runtime.us-west-2.vpce.amazonaws.com" + ) + # runtime 端点填给控制面服务名时不匹配标准形态,同样透传 + assert ( + resolve("bedrock", "https://bedrock-runtime.us-east-1.amazonaws.com") + == "https://bedrock-runtime.us-east-1.amazonaws.com" + ) + + +def test_create_bedrock_client_uses_custom_endpoint(): + """创建 Bedrock 客户端时应把 PrivateLink 地址传给 boto3""" + manager = LLMProviderManager() + endpoint_url = ( + "https://vpce-0abc123-xyz.bedrock-runtime.us-west-2.vpce.amazonaws.com" + ) + client = MagicMock() + + with patch("boto3.client", return_value=client) as create_client: + result = manager.create_bedrock_client( + service_name="bedrock-runtime", + region="us-west-2", + credentials={ + "auth_scheme": "sigv4", + "access_key_id": "AKIAIOSFODNN7EXAMPLE", + "secret_access_key": "secret", + }, + base_url=endpoint_url, + use_proxy=False, + ) + + assert result is client + assert create_client.call_args.kwargs["endpoint_url"] == endpoint_url + + +def test_resolve_runtime_bedrock_bearer(): + manager = LLMProviderManager() + runtime = asyncio.run( + manager.resolve_runtime( + provider_id="amazon-bedrock", + model="global.anthropic.claude-haiku-4-5-20251001-v1:0", + api_key="bedrock-api-key-abc123", + base_url="https://bedrock-runtime.ap-northeast-1.amazonaws.com", + ) + ) + + assert runtime["runtime"] == "bedrock" + assert runtime["aws_region"] == "ap-northeast-1" + assert runtime["aws_auth"]["auth_scheme"] == "bearer" + + +def test_resolve_runtime_bedrock_sigv4_default_region(): + manager = LLMProviderManager() + runtime = asyncio.run( + manager.resolve_runtime( + provider_id="amazon-bedrock", + model="us.anthropic.claude-sonnet-4-5-20250929-v1:0", + api_key="AKIAIOSFODNN7EXAMPLE:wJalrXUtnFEMI/K7MDENG", + ) + ) + + assert runtime["runtime"] == "bedrock" + assert runtime["aws_region"] == "us-east-1" + assert runtime["aws_auth"]["auth_scheme"] == "sigv4" + assert runtime["aws_auth"]["access_key_id"] == "AKIAIOSFODNN7EXAMPLE" + + +def test_resolve_runtime_bedrock_missing_credentials_rejected(): + manager = LLMProviderManager() + with pytest.raises(LLMProviderAuthError): + asyncio.run( + manager.resolve_runtime( + provider_id="amazon-bedrock", + model="us.anthropic.claude-sonnet-4-5-20250929-v1:0", + api_key=None, + ) + ) + + +def test_bedrock_model_matches_region(): + """目录模型应按 Profile 分区及裸模型 ON_DEMAND Region 过滤""" + matches = LLMProviderManager._bedrock_model_matches_region + + # 已知裸模型 ID 仅在其支持 ON_DEMAND 的 Region 保留 + assert matches("anthropic.claude-3-5-sonnet-20241022-v2:0", "us-west-2") + assert matches("anthropic.claude-3-5-sonnet-20241022-v2:0", "ap-southeast-2") + assert not matches("anthropic.claude-3-5-sonnet-20241022-v2:0", "ap-northeast-1") + assert not matches("anthropic.claude-sonnet-4-5-20250929-v1:0", "us-west-2") + assert not matches("amazon.nova-premier-v1:0", "ap-northeast-1") + assert not matches("meta.llama4-maverick-17b-instruct-v1:0", "ap-northeast-1") + # 已确认支持 ON_DEMAND 的裸模型与 global Profile 维持可用 + assert matches("amazon.nova-lite-v1:0", "ap-northeast-1") + assert matches("openai.gpt-oss-20b-1:0", "ap-northeast-1") + assert not matches("openai.gpt-oss-20b-1:0", "ap-southeast-1") + assert matches("global.anthropic.claude-sonnet-4-5-20250929-v1:0", "ap-northeast-1") + assert not matches( + "global.anthropic.claude-sonnet-4-5-20250929-v1:0", + "us-gov-west-1", + ) + # 地理前缀只在对应分区 Region 可调用 + assert matches("us.anthropic.claude-haiku-4-5-20251001-v1:0", "us-west-2") + assert not matches("us.anthropic.claude-haiku-4-5-20251001-v1:0", "ap-northeast-1") + assert not matches("us.anthropic.claude-haiku-4-5-20251001-v1:0", "us-gov-west-1") + assert matches("apac.amazon.nova-micro-v1:0", "ap-southeast-1") + assert not matches("apac.amazon.nova-micro-v1:0", "eu-central-1") + assert matches("eu.anthropic.claude-haiku-4-5-20251001-v1:0", "eu-central-1") + assert not matches("eu.anthropic.claude-haiku-4-5-20251001-v1:0", "us-east-1") + assert not matches("eu.anthropic.claude-haiku-4-5-20251001-v1:0", "eu-isoe-west-1") + + +def test_bedrock_au_profile_matches_melbourne_region(): + """AU Inference Profile 应允许从悉尼和墨尔本 Region 调用""" + matches = LLMProviderManager._bedrock_model_matches_region + + assert matches("au.amazon.nova-lite-v1:0", "ap-southeast-2") + assert matches("au.amazon.nova-lite-v1:0", "ap-southeast-4") + + +def test_list_models_bedrock_custom_endpoint_skips_control_plane(): + """自定义 runtime 端点刷新模型时应直接使用离线目录""" + manager = LLMProviderManager() + manager._models_dev_data = { + "amazon-bedrock": { + "id": "amazon-bedrock", + "name": "Amazon Bedrock", + "models": { + "global.anthropic.claude-sonnet-4-5-20250929-v1:0": { + "name": "Claude Sonnet 4.5 (Global)", + "limit": {"context": 200000, "output": 64000}, + }, + }, + } + } + manager._models_dev_loaded_at = time.time() + + with patch.object( + LLMProviderManager, + "create_bedrock_client", + side_effect=AssertionError("不应访问控制面"), + ): + models = asyncio.run( + manager._list_models_from_bedrock( + api_key="bedrock-api-key-runtime-only", + base_url="https://bedrock-runtime-fips.us-east-1.amazonaws.com", + use_proxy=False, + ) + ) + + assert [model["id"] for model in models] == [ + "global.anthropic.claude-sonnet-4-5-20250929-v1:0" + ] + + +def test_list_models_bedrock_keeps_control_plane_on_demand_models(): + """控制面返回的 ON_DEMAND 基础模型不应被静态降级规则遗漏""" + manager = LLMProviderManager() + client = MagicMock() + client.get_paginator.return_value.paginate.return_value = [ + { + "inferenceProfileSummaries": [ + { + "inferenceProfileId": ( + "global.anthropic.claude-sonnet-4-5-20250929-v1:0" + ), + "inferenceProfileName": "Claude Sonnet 4.5 (Global)", + "status": "ACTIVE", + } + ] + } + ] + client.list_foundation_models.return_value = { + "modelSummaries": [ + { + "modelId": "openai.gpt-oss-20b-1:0", + "modelName": "GPT OSS 20B", + "modelLifecycle": {"status": "ACTIVE"}, + }, + { + "modelId": "amazon.nova-lite-v1:0", + "modelName": "Nova Lite", + "modelLifecycle": {"status": "ACTIVE"}, + }, + ] + } + + with patch.object( + LLMProviderManager, "create_bedrock_client", return_value=client + ): + models = asyncio.run( + manager._list_models_from_bedrock( + api_key="bedrock-api-key-runtime-only", + base_url="https://bedrock-runtime.ap-northeast-1.amazonaws.com", + use_proxy=False, + ) + ) + + assert {model["id"] for model in models} == { + "amazon.nova-lite-v1:0", + "global.anthropic.claude-sonnet-4-5-20250929-v1:0", + "openai.gpt-oss-20b-1:0", + } + client.close.assert_called_once() + + +def test_list_models_bedrock_falls_back_to_models_dev_on_control_plane_denial(): + """控制面被拒(如 API Key 仅授权 bedrock-runtime)时降级 models.dev 目录""" + manager = LLMProviderManager() + # 预填 models.dev 内存缓存,降级路径不触发真实网络请求 + manager._models_dev_data = { + "amazon-bedrock": { + "id": "amazon-bedrock", + "name": "Amazon Bedrock", + "models": { + "anthropic.claude-3-5-sonnet-20241022-v2:0": { + "name": "Claude Sonnet 3.5 v2", + "limit": {"context": 200000, "output": 8192}, + }, + "amazon.nova-lite-v1:0": { + "name": "Nova Lite", + "limit": {"context": 300000, "output": 5000}, + }, + "openai.gpt-oss-20b-1:0": { + "name": "GPT OSS 20B", + "limit": {"context": 131072, "output": 16384}, + }, + "apac.amazon.nova-lite-v1:0": { + "name": "Nova Lite (APAC)", + "limit": {"context": 300000, "output": 5000}, + }, + "meta.llama4-maverick-17b-instruct-v1:0": { + "name": "Llama 4 Maverick", + "limit": {"context": 1000000, "output": 8192}, + }, + "anthropic.claude-sonnet-4-5-20250929-v1:0": { + "name": "Claude Sonnet 4.5", + "limit": {"context": 200000, "output": 64000}, + }, + "global.anthropic.claude-sonnet-4-5-20250929-v1:0": { + "name": "Claude Sonnet 4.5 (Global)", + "limit": {"context": 200000, "output": 64000}, + }, + "us.anthropic.claude-haiku-4-5-20251001-v1:0": { + "name": "Claude Haiku 4.5 (US)", + "limit": {"context": 200000, "output": 64000}, + }, + }, + } + } + manager._models_dev_loaded_at = time.time() + + denied_client = MagicMock() + denied_client.get_paginator.side_effect = Exception( + "AccessDeniedException: not authorized to perform bedrock:ListInferenceProfiles" + ) + + with patch.object( + LLMProviderManager, "create_bedrock_client", return_value=denied_client + ): + models = asyncio.run( + manager._list_models_from_bedrock( + api_key="bedrock-api-key-runtime-only", + base_url="https://bedrock-runtime.ap-northeast-1.amazonaws.com", + use_proxy=False, + ) + ) + + # 降级后仅保留东京 Region 可调用的裸模型与 Profile + model_ids = {m["id"] for m in models} + assert "anthropic.claude-3-5-sonnet-20241022-v2:0" not in model_ids + assert "amazon.nova-lite-v1:0" in model_ids + assert "openai.gpt-oss-20b-1:0" in model_ids + assert "apac.amazon.nova-lite-v1:0" in model_ids + assert "meta.llama4-maverick-17b-instruct-v1:0" not in model_ids + assert "global.anthropic.claude-sonnet-4-5-20250929-v1:0" in model_ids + assert "us.anthropic.claude-haiku-4-5-20251001-v1:0" not in model_ids + assert "anthropic.claude-sonnet-4-5-20250929-v1:0" not in model_ids + assert all(m["source"] == "models.dev" for m in models) + denied_client.close.assert_called_once()