358 lines
14 KiB
Python
358 lines
14 KiB
Python
"""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 (
|
|
AgentConfigurationError,
|
|
InvalidStateTransition,
|
|
PermissionDenied,
|
|
ProviderConnectionError,
|
|
)
|
|
from common.logging import sanitize_summary
|
|
|
|
from .models import (
|
|
AgentRun,
|
|
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.WAITING_APPROVAL: {RunStatus.RUNNING, RunStatus.FAILED, RunStatus.CANCELLED},
|
|
}
|
|
|
|
|
|
def runs_for_user(user):
|
|
"""普通站点始终限制为当前用户数据,管理员跨用户查询使用 Admin。"""
|
|
|
|
return AgentRun.objects.filter(owner=user)
|
|
|
|
|
|
def _append_event(run, event_type: str, summary: str, payload=None):
|
|
"""在持有 Run 写锁的事务中分配下一个事件序号。"""
|
|
|
|
last = run.events.aggregate(value=Max("sequence"))["value"] or 0
|
|
return AgentRunEvent.objects.create(
|
|
run=run,
|
|
sequence=last + 1,
|
|
event_type=event_type,
|
|
summary=summary,
|
|
payload_summary=sanitize_summary(payload or {}),
|
|
occurred_at=timezone.now(),
|
|
)
|
|
|
|
|
|
@transaction.atomic
|
|
def create_run(owner, title: str, input_summary=None) -> AgentRun:
|
|
"""创建等待运行的记录,并原子追加创建事件。"""
|
|
|
|
run = AgentRun.objects.create(
|
|
owner=owner, title=title.strip(), input_summary=sanitize_summary(input_summary or {})
|
|
)
|
|
_append_event(run, "run_created", "运行记录已创建")
|
|
return run
|
|
|
|
|
|
@transaction.atomic
|
|
def transition_run(run_id, target_status: str, *, error_code="", error_summary="", output=None):
|
|
"""校验并执行一次状态转换,主记录与审计事件同时提交。"""
|
|
|
|
run = AgentRun.objects.select_for_update().get(pk=run_id)
|
|
if target_status not in TRANSITIONS.get(run.status, set()):
|
|
raise InvalidStateTransition(f"不允许从 {run.status} 转换到 {target_status}。")
|
|
now = timezone.now()
|
|
run.status = target_status
|
|
run.lock_version += 1
|
|
if target_status == RunStatus.RUNNING and run.started_at is None:
|
|
run.started_at = now
|
|
if target_status in {RunStatus.SUCCEEDED, RunStatus.FAILED, RunStatus.CANCELLED}:
|
|
run.finished_at = now
|
|
if run.started_at:
|
|
run.duration_ms = max(0, int((now - run.started_at).total_seconds() * 1000))
|
|
if target_status == RunStatus.SUCCEEDED:
|
|
run.output_summary = sanitize_summary(output or {})
|
|
if target_status == RunStatus.FAILED:
|
|
run.error_code = error_code[:80]
|
|
run.error_summary = error_summary[:2000]
|
|
run.save()
|
|
_append_event(run, f"run_{target_status}", f"运行状态变更为 {run.get_status_display()}")
|
|
return run
|
|
|
|
|
|
@transaction.atomic
|
|
def request_approval(run_id, request_key: str, approval_type: str, summary=None):
|
|
"""暂停运行并创建唯一的待确认请求。"""
|
|
|
|
run = AgentRun.objects.select_for_update().get(pk=run_id)
|
|
if run.status != RunStatus.RUNNING:
|
|
raise InvalidStateTransition("只有运行中的任务可以请求人工确认。")
|
|
approval = HumanApproval.objects.create(
|
|
run=run,
|
|
request_key=request_key,
|
|
approval_type=approval_type,
|
|
request_summary=sanitize_summary(summary or {}),
|
|
requested_at=timezone.now(),
|
|
)
|
|
run.status = RunStatus.WAITING_APPROVAL
|
|
run.lock_version += 1
|
|
run.save(update_fields=("status", "lock_version", "updated_at"))
|
|
_append_event(run, "approval_requested", "运行等待人工确认", {"request_key": request_key})
|
|
return approval
|
|
|
|
|
|
@transaction.atomic
|
|
def start_tool_call(run_id, call_id: str, tool_name: str, arguments=None, idempotency_key=""):
|
|
"""在外部调用前登记开始状态;唯一约束负责阻止重复调用标识和幂等键。"""
|
|
|
|
run = AgentRun.objects.select_for_update().get(pk=run_id)
|
|
if run.status != RunStatus.RUNNING:
|
|
raise InvalidStateTransition("只有运行中的任务可以开始工具调用。")
|
|
call = ToolCall.objects.create(
|
|
run=run,
|
|
call_id=call_id,
|
|
tool_name=tool_name,
|
|
idempotency_key=idempotency_key,
|
|
arguments_summary=sanitize_summary(arguments or {}),
|
|
started_at=timezone.now(),
|
|
)
|
|
_append_event(run, "tool_started", f"工具 {tool_name} 开始执行", {"call_id": call_id})
|
|
return call
|
|
|
|
|
|
@transaction.atomic
|
|
def finish_tool_call(call_id: int, *, result=None, error_code="", error_summary=""):
|
|
"""在外部调用结束后的独立短事务中登记成功或失败结果。"""
|
|
|
|
call = ToolCall.objects.select_for_update().select_related("run").get(pk=call_id)
|
|
if call.status != ToolCallStatus.STARTED:
|
|
raise InvalidStateTransition("工具调用已经结束。")
|
|
now = timezone.now()
|
|
call.finished_at = now
|
|
call.duration_ms = max(0, int((now - call.started_at).total_seconds() * 1000))
|
|
if error_code:
|
|
call.status = ToolCallStatus.FAILED
|
|
call.error_code = error_code[:80]
|
|
call.error_summary = error_summary[:2000]
|
|
event_type, summary = "tool_failed", f"工具 {call.tool_name} 执行失败"
|
|
else:
|
|
call.status = ToolCallStatus.SUCCEEDED
|
|
call.result_summary = sanitize_summary(result or {})
|
|
event_type, summary = "tool_completed", f"工具 {call.tool_name} 执行完成"
|
|
call.save()
|
|
run = AgentRun.objects.select_for_update().get(pk=call.run_id)
|
|
_append_event(run, event_type, summary, {"call_id": call.call_id})
|
|
return call
|
|
|
|
|
|
@transaction.atomic
|
|
def resolve_approval(actor, approval_id, approved: bool, summary=None):
|
|
"""只允许所属用户或管理员处理一次待确认请求。"""
|
|
|
|
approval = HumanApproval.objects.select_for_update().select_related("run").get(pk=approval_id)
|
|
if actor != approval.run.owner and not actor.is_staff:
|
|
raise PermissionDenied("无权处理该确认请求。")
|
|
if approval.status != ApprovalStatus.PENDING:
|
|
raise InvalidStateTransition("该确认请求已经处理。")
|
|
approval.status = ApprovalStatus.APPROVED if approved else ApprovalStatus.REJECTED
|
|
approval.resolved_by = actor
|
|
approval.resolved_at = timezone.now()
|
|
approval.decision_summary = sanitize_summary(summary or {})
|
|
approval.save()
|
|
run = AgentRun.objects.select_for_update().get(pk=approval.run_id)
|
|
run.status = RunStatus.RUNNING if approved else RunStatus.CANCELLED
|
|
run.lock_version += 1
|
|
if not approved:
|
|
run.finished_at = timezone.now()
|
|
run.save()
|
|
_append_event(run, "approval_resolved", "人工确认已处理", {"approved": approved})
|
|
return approval
|