feat(agent-runtime): 接入多服务商安全配置
This commit is contained in:
+190
-3
@@ -1,10 +1,20 @@
|
||||
"""Agent Run 的事务、权限和状态转换服务。"""
|
||||
"""Agent Run 与模型服务商配置的领域服务。"""
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.request import Request, urlopen
|
||||
|
||||
from django.db import transaction
|
||||
from django.db.models import Max
|
||||
from django.utils import timezone
|
||||
|
||||
from common.exceptions import InvalidStateTransition, PermissionDenied
|
||||
from common.exceptions import (
|
||||
AgentConfigurationError,
|
||||
InvalidStateTransition,
|
||||
PermissionDenied,
|
||||
ProviderConnectionError,
|
||||
)
|
||||
from common.logging import sanitize_summary
|
||||
|
||||
from .models import (
|
||||
@@ -12,14 +22,191 @@ from .models import (
|
||||
AgentRunEvent,
|
||||
ApprovalStatus,
|
||||
HumanApproval,
|
||||
ModelProviderConfig,
|
||||
ProviderConnectionStatus,
|
||||
RunStatus,
|
||||
ToolCall,
|
||||
ToolCallStatus,
|
||||
)
|
||||
from .providers import ModelProvider, get_provider
|
||||
from .secrets import decrypt_api_key, encrypt_api_key
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProviderVerificationResult:
|
||||
"""连接验证后的安全结果,不携带凭据或上游响应正文。"""
|
||||
|
||||
models: tuple[str, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RuntimeModelConfig:
|
||||
"""真实 Gateway 调用前解析出的用户级模型配置。"""
|
||||
|
||||
provider_code: str
|
||||
protocol: str
|
||||
base_url: str
|
||||
api_key: str
|
||||
model: str
|
||||
|
||||
|
||||
def provider_configs_for_user(user):
|
||||
"""返回当前用户配置,并由调用方按服务商编码组织展示。"""
|
||||
|
||||
return ModelProviderConfig.objects.filter(owner=user)
|
||||
|
||||
|
||||
def _verification_headers(provider: ModelProvider, api_key: str) -> dict[str, str]:
|
||||
"""按预设协议构造最小认证头;API Key 不进入 URL。"""
|
||||
|
||||
headers = {"Accept": "application/json", "User-Agent": "JobRadar/0.1"}
|
||||
if provider.protocol == "anthropic_messages":
|
||||
headers.update({"x-api-key": api_key, "anthropic-version": "2023-06-01"})
|
||||
elif provider.protocol == "gemini":
|
||||
headers["x-goog-api-key"] = api_key
|
||||
else:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
return headers
|
||||
|
||||
|
||||
def verify_provider_api_key(
|
||||
provider: ModelProvider, api_key: str, *, opener=None
|
||||
) -> ProviderVerificationResult:
|
||||
"""通过模型目录验证凭据,禁止执行会产生模型用量的对话请求。"""
|
||||
|
||||
if not provider.key_only or not provider.base_url:
|
||||
raise ProviderConnectionError("extra_configuration_required", "该服务商需要额外配置。")
|
||||
open_request = opener or urlopen
|
||||
endpoint = f"{provider.base_url.rstrip('/')}/{provider.models_path.lstrip('/')}"
|
||||
request = Request(endpoint, headers=_verification_headers(provider, api_key), method="GET")
|
||||
try:
|
||||
with open_request(request, timeout=12) as response:
|
||||
payload = json.loads(response.read().decode("utf-8"))
|
||||
except HTTPError as exc:
|
||||
if exc.code == 401:
|
||||
raise ProviderConnectionError("invalid_api_key", "API Key 无效或已失效。") from exc
|
||||
if exc.code == 403:
|
||||
raise ProviderConnectionError(
|
||||
"permission_denied", "API Key 没有访问模型目录的权限。"
|
||||
) from exc
|
||||
if exc.code == 429:
|
||||
raise ProviderConnectionError(
|
||||
"rate_limited", "服务商请求过于频繁,请稍后重试。"
|
||||
) from exc
|
||||
raise ProviderConnectionError("provider_http_error", "服务商暂时无法完成验证。") from exc
|
||||
except (URLError, TimeoutError) as exc:
|
||||
raise ProviderConnectionError(
|
||||
"provider_unreachable", "无法连接服务商,请稍后重试。"
|
||||
) from exc
|
||||
except (json.JSONDecodeError, UnicodeError) as exc:
|
||||
raise ProviderConnectionError(
|
||||
"invalid_provider_response", "服务商返回了无法识别的数据。"
|
||||
) from exc
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
raise ProviderConnectionError(
|
||||
"invalid_provider_response", "服务商返回了无法识别的数据。"
|
||||
)
|
||||
raw_models = payload.get("data") or payload.get("models") or []
|
||||
model_ids = []
|
||||
for item in raw_models:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
model_id = item.get("id") or item.get("name")
|
||||
if isinstance(model_id, str) and model_id:
|
||||
model_ids.append(model_id.removeprefix("models/"))
|
||||
# 某些兼容端点只验证凭据但不返回模型目录,保留经过项目核验的推荐模型。
|
||||
if not model_ids and provider.default_model:
|
||||
model_ids.append(provider.default_model)
|
||||
return ProviderVerificationResult(tuple(dict.fromkeys(model_ids)))
|
||||
|
||||
|
||||
@transaction.atomic
|
||||
def save_verified_provider_config(
|
||||
owner,
|
||||
provider: ModelProvider,
|
||||
api_key: str,
|
||||
result: ProviderVerificationResult,
|
||||
*,
|
||||
model_id: str = "",
|
||||
use_as_default: bool = False,
|
||||
) -> ModelProviderConfig:
|
||||
"""只在连接验证成功后替换密钥,避免错误输入破坏已有可用配置。"""
|
||||
|
||||
selected_model = model_id.strip() or provider.default_model
|
||||
if not selected_model and result.models:
|
||||
selected_model = result.models[0]
|
||||
existing_count = ModelProviderConfig.objects.filter(owner=owner).count()
|
||||
should_default = use_as_default or existing_count == 0
|
||||
if should_default:
|
||||
ModelProviderConfig.objects.filter(owner=owner, is_default=True).update(is_default=False)
|
||||
config, _ = ModelProviderConfig.objects.update_or_create(
|
||||
owner=owner,
|
||||
provider_code=provider.code,
|
||||
defaults={
|
||||
"encrypted_api_key": encrypt_api_key(api_key),
|
||||
"key_hint": f"••••{api_key[-4:]}",
|
||||
"status": ProviderConnectionStatus.CONNECTED,
|
||||
"enabled": True,
|
||||
"is_default": should_default,
|
||||
"default_model_id": selected_model,
|
||||
"available_models": list(result.models)[:200],
|
||||
"last_verified_at": timezone.now(),
|
||||
"last_error_code": "",
|
||||
},
|
||||
)
|
||||
return config
|
||||
|
||||
|
||||
@transaction.atomic
|
||||
def disconnect_provider(owner, provider_code: str) -> bool:
|
||||
"""删除当前用户指定服务商的密文配置,不影响其他用户或运行审计数据。"""
|
||||
|
||||
deleted, _ = ModelProviderConfig.objects.filter(
|
||||
owner=owner, provider_code=provider_code
|
||||
).delete()
|
||||
if deleted and not ModelProviderConfig.objects.filter(owner=owner, is_default=True).exists():
|
||||
fallback = ModelProviderConfig.objects.filter(owner=owner, enabled=True).first()
|
||||
if fallback:
|
||||
fallback.is_default = True
|
||||
fallback.save(update_fields=("is_default", "updated_at"))
|
||||
return bool(deleted)
|
||||
|
||||
|
||||
def resolve_runtime_model_config(owner, provider_code: str = "") -> RuntimeModelConfig:
|
||||
"""在真实执行边界解密当前用户配置,密钥不会进入运行请求或数据库摘要。"""
|
||||
|
||||
queryset = ModelProviderConfig.objects.filter(
|
||||
owner=owner,
|
||||
enabled=True,
|
||||
status=ProviderConnectionStatus.CONNECTED,
|
||||
)
|
||||
config = queryset.filter(provider_code=provider_code).first() if provider_code else None
|
||||
config = config or queryset.filter(is_default=True).first()
|
||||
if config is None:
|
||||
raise AgentConfigurationError("当前用户尚未配置可用的模型服务商。")
|
||||
provider = get_provider(config.provider_code)
|
||||
if provider is None or not provider.base_url:
|
||||
raise AgentConfigurationError("当前模型服务商预设不可用,请重新配置。")
|
||||
if not config.default_model_id:
|
||||
raise AgentConfigurationError("当前服务商尚未选择默认模型。")
|
||||
return RuntimeModelConfig(
|
||||
provider_code=provider.code,
|
||||
protocol=provider.protocol,
|
||||
base_url=provider.base_url,
|
||||
api_key=decrypt_api_key(config.encrypted_api_key),
|
||||
model=config.default_model_id,
|
||||
)
|
||||
|
||||
|
||||
TRANSITIONS = {
|
||||
RunStatus.PENDING: {RunStatus.RUNNING, RunStatus.CANCELLED},
|
||||
RunStatus.RUNNING: {RunStatus.WAITING_APPROVAL, RunStatus.SUCCEEDED, RunStatus.FAILED, RunStatus.CANCELLED},
|
||||
RunStatus.RUNNING: {
|
||||
RunStatus.WAITING_APPROVAL,
|
||||
RunStatus.SUCCEEDED,
|
||||
RunStatus.FAILED,
|
||||
RunStatus.CANCELLED,
|
||||
},
|
||||
RunStatus.WAITING_APPROVAL: {RunStatus.RUNNING, RunStatus.FAILED, RunStatus.CANCELLED},
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user