307 lines
14 KiB
Python
307 lines
14 KiB
Python
from unittest.mock import patch
|
|
|
|
from django.contrib.auth import get_user_model
|
|
from django.db import IntegrityError, transaction
|
|
from django.test import SimpleTestCase, TestCase, override_settings
|
|
from django.urls import reverse
|
|
|
|
from common.exceptions import (
|
|
AgentConfigurationError,
|
|
InvalidStateTransition,
|
|
PermissionDenied,
|
|
ProviderConnectionError,
|
|
)
|
|
|
|
from .config import AgentRuntimeConfig
|
|
from .gateway import AgentExecutionResult, StubAgentRunnerGateway
|
|
from .models import ApprovalStatus, ModelProviderConfig, RunStatus, ToolCallStatus
|
|
from .orchestration import execute_run
|
|
from .providers import get_provider
|
|
from .secrets import decrypt_api_key, encrypt_api_key
|
|
from .services import (
|
|
ProviderVerificationResult,
|
|
create_run,
|
|
disconnect_provider,
|
|
finish_tool_call,
|
|
request_approval,
|
|
resolve_approval,
|
|
resolve_runtime_model_config,
|
|
save_verified_provider_config,
|
|
start_tool_call,
|
|
transition_run,
|
|
verify_provider_api_key,
|
|
)
|
|
|
|
|
|
class ExplodingGateway:
|
|
"""模拟外部边界抛错,并故意在异常文本中携带不应落库的秘密。"""
|
|
|
|
def run(self, request):
|
|
raise RuntimeError("api_key=should-not-be-stored")
|
|
|
|
|
|
class AgentRunTests(TestCase):
|
|
"""覆盖运行状态、事件顺序、脱敏和跨用户拒绝路径。"""
|
|
|
|
def setUp(self):
|
|
users = get_user_model().objects
|
|
self.alice = users.create_user(username="alice", password="safe-pass-123")
|
|
self.bob = users.create_user(username="bob", password="safe-pass-123")
|
|
|
|
def test_run_lifecycle_and_sensitive_summary(self):
|
|
run = create_run(self.alice, "基线运行", {"api_key": "secret", "query": "Python"})
|
|
self.assertEqual(run.input_summary["api_key"], "***")
|
|
transition_run(run.id, RunStatus.RUNNING)
|
|
finished = transition_run(run.id, RunStatus.SUCCEEDED, output={"count": 1})
|
|
self.assertEqual(finished.status, RunStatus.SUCCEEDED)
|
|
self.assertEqual(list(finished.events.values_list("sequence", flat=True)), [1, 2, 3])
|
|
with self.assertRaises(InvalidStateTransition):
|
|
transition_run(run.id, RunStatus.RUNNING)
|
|
|
|
def test_cross_user_detail_returns_404(self):
|
|
run = create_run(self.alice, "私有运行")
|
|
self.client.force_login(self.bob)
|
|
response = self.client.get(reverse("agent_runtime:run-detail", args=(run.id,)))
|
|
self.assertEqual(response.status_code, 404)
|
|
|
|
def test_run_list_only_contains_current_user_data(self):
|
|
create_run(self.alice, "爱丽丝的运行")
|
|
create_run(self.bob, "鲍勃的运行")
|
|
self.client.force_login(self.alice)
|
|
|
|
response = self.client.get(reverse("agent_runtime:run-list"))
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertContains(response, "爱丽丝的运行")
|
|
self.assertNotContains(response, "鲍勃的运行")
|
|
|
|
def test_tool_call_is_recorded_and_sanitized(self):
|
|
run = create_run(self.alice, "工具运行")
|
|
transition_run(run.id, RunStatus.RUNNING)
|
|
call = start_tool_call(run.id, "call-1", "demo_tool", {"password": "secret"})
|
|
self.assertEqual(call.arguments_summary["password"], "***")
|
|
finished = finish_tool_call(call.id, result={"count": 1})
|
|
self.assertEqual(finished.status, "succeeded")
|
|
self.assertEqual(finished.result_summary, {"count": 1})
|
|
|
|
def test_tool_call_failure_and_duplicate_finish_are_recorded(self):
|
|
run = create_run(self.alice, "失败工具运行")
|
|
transition_run(run.id, RunStatus.RUNNING)
|
|
call = start_tool_call(run.id, "call-1", "demo_tool", idempotency_key="same-operation")
|
|
failed = finish_tool_call(call.id, error_code="timeout", error_summary="请求超时")
|
|
self.assertEqual(failed.status, ToolCallStatus.FAILED)
|
|
self.assertEqual(failed.error_code, "timeout")
|
|
with self.assertRaises(InvalidStateTransition):
|
|
finish_tool_call(call.id, result={"unexpected": True})
|
|
|
|
def test_duplicate_tool_idempotency_key_is_rejected(self):
|
|
run = create_run(self.alice, "幂等工具运行")
|
|
transition_run(run.id, RunStatus.RUNNING)
|
|
start_tool_call(run.id, "call-1", "demo_tool", idempotency_key="same-operation")
|
|
with self.assertRaises(IntegrityError), transaction.atomic():
|
|
start_tool_call(run.id, "call-2", "demo_tool", idempotency_key="same-operation")
|
|
|
|
def test_approval_can_be_approved_once_by_owner(self):
|
|
run = create_run(self.alice, "确认运行")
|
|
transition_run(run.id, RunStatus.RUNNING)
|
|
approval = request_approval(run.id, "approval-1", "external_action", {"token": "secret"})
|
|
self.assertEqual(approval.request_summary["token"], "***")
|
|
resolved = resolve_approval(self.alice, approval.id, True, {"reason": "允许"})
|
|
self.assertEqual(resolved.status, ApprovalStatus.APPROVED)
|
|
run.refresh_from_db()
|
|
self.assertEqual(run.status, RunStatus.RUNNING)
|
|
with self.assertRaises(InvalidStateTransition):
|
|
resolve_approval(self.alice, approval.id, True)
|
|
|
|
def test_approval_rejects_cross_user_and_can_cancel_run(self):
|
|
run = create_run(self.alice, "拒绝确认运行")
|
|
transition_run(run.id, RunStatus.RUNNING)
|
|
approval = request_approval(run.id, "approval-1", "external_action")
|
|
with self.assertRaises(PermissionDenied):
|
|
resolve_approval(self.bob, approval.id, False)
|
|
resolve_approval(self.alice, approval.id, False)
|
|
run.refresh_from_db()
|
|
self.assertEqual(run.status, RunStatus.CANCELLED)
|
|
|
|
def test_stub_gateway_drives_successful_persistent_run(self):
|
|
run = create_run(self.alice, "Stub 成功运行", {"query": "Python"})
|
|
gateway = StubAgentRunnerGateway(AgentExecutionResult(True, {"count": 2}))
|
|
finished = execute_run(run.id, gateway)
|
|
self.assertEqual(finished.status, RunStatus.SUCCEEDED)
|
|
self.assertEqual(finished.output_summary, {"count": 2})
|
|
self.assertEqual(finished.events.count(), 3)
|
|
|
|
def test_stub_gateway_failure_and_exception_reach_failed_state(self):
|
|
failed_run = create_run(self.alice, "Stub 失败运行")
|
|
failed_gateway = StubAgentRunnerGateway(
|
|
AgentExecutionResult(False, error_code="model_error", error_summary="模型失败")
|
|
)
|
|
failed = execute_run(failed_run.id, failed_gateway)
|
|
self.assertEqual(failed.status, RunStatus.FAILED)
|
|
self.assertEqual(failed.error_code, "model_error")
|
|
|
|
exploding_run = create_run(self.alice, "Stub 异常运行")
|
|
exploded = execute_run(exploding_run.id, ExplodingGateway())
|
|
self.assertEqual(exploded.status, RunStatus.FAILED)
|
|
self.assertNotIn("should-not-be-stored", exploded.error_summary)
|
|
|
|
|
|
class AgentRuntimeConfigTests(SimpleTestCase):
|
|
"""验证真实执行配置只在调用边界检查,不影响其他 Django 功能。"""
|
|
|
|
@override_settings(OPENAI_API_KEY="", OPENAI_MODEL="", AGENT_DEFAULT_NAME="job_research")
|
|
def test_real_execution_requires_key_and_model(self):
|
|
with self.assertRaises(AgentConfigurationError):
|
|
AgentRuntimeConfig.from_settings().require_real_execution()
|
|
|
|
@override_settings(
|
|
OPENAI_API_KEY="test-key", OPENAI_MODEL="test-model", AGENT_DEFAULT_NAME="job_research"
|
|
)
|
|
def test_complete_configuration_is_accepted_without_external_call(self):
|
|
config = AgentRuntimeConfig.from_settings().require_real_execution()
|
|
self.assertEqual(config.model, "test-model")
|
|
|
|
|
|
class FakeModelsResponse:
|
|
"""模拟只返回公开模型标识的服务商模型目录。"""
|
|
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc_value, traceback):
|
|
return False
|
|
|
|
def read(self):
|
|
return b'{"data":[{"id":"model-a"},{"id":"model-b"}]}'
|
|
|
|
|
|
class ProviderVerificationTests(SimpleTestCase):
|
|
"""验证连接测试只访问模型目录,且凭据不会进入 URL。"""
|
|
|
|
def test_openai_compatible_verification_uses_bearer_header(self):
|
|
captured = {}
|
|
|
|
def opener(request, timeout):
|
|
captured["url"] = request.full_url
|
|
captured["authorization"] = request.get_header("Authorization")
|
|
captured["timeout"] = timeout
|
|
return FakeModelsResponse()
|
|
|
|
result = verify_provider_api_key(get_provider("deepseek"), "safe-test-key", opener=opener)
|
|
|
|
self.assertEqual(result.models, ("model-a", "model-b"))
|
|
self.assertEqual(captured["authorization"], "Bearer safe-test-key")
|
|
self.assertNotIn("safe-test-key", captured["url"])
|
|
self.assertEqual(captured["timeout"], 12)
|
|
|
|
def test_provider_requiring_extra_configuration_is_rejected_without_request(self):
|
|
with self.assertRaises(ProviderConnectionError) as context:
|
|
verify_provider_api_key(get_provider("azure_openai"), "safe-test-key")
|
|
self.assertEqual(context.exception.code, "extra_configuration_required")
|
|
|
|
|
|
@override_settings(MODEL_API_KEY_ENCRYPTION_KEY="test-only-model-key-32-characters-minimum")
|
|
class ModelProviderConfigTests(TestCase):
|
|
"""覆盖密钥加密、用户隔离、默认选择和页面不回显边界。"""
|
|
|
|
def setUp(self):
|
|
users = get_user_model().objects
|
|
self.alice = users.create_user(username="alice", password="safe-pass-123")
|
|
self.bob = users.create_user(username="bob", password="safe-pass-123")
|
|
self.provider = get_provider("openai")
|
|
self.result = ProviderVerificationResult(("model-a", "model-b"))
|
|
|
|
def test_api_key_round_trip_uses_ciphertext(self):
|
|
ciphertext = encrypt_api_key("safe-test-key")
|
|
self.assertNotIn("safe-test-key", ciphertext)
|
|
self.assertEqual(decrypt_api_key(ciphertext), "safe-test-key")
|
|
|
|
def test_verified_config_is_user_owned_and_resolves_for_runtime(self):
|
|
config = save_verified_provider_config(
|
|
self.alice,
|
|
self.provider,
|
|
"safe-test-key",
|
|
self.result,
|
|
model_id="model-a",
|
|
)
|
|
self.assertTrue(config.is_default)
|
|
runtime = resolve_runtime_model_config(self.alice)
|
|
self.assertEqual(runtime.api_key, "safe-test-key")
|
|
self.assertEqual(runtime.model, "model-a")
|
|
user_config = AgentRuntimeConfig.from_user(self.alice)
|
|
self.assertEqual(user_config.provider_code, "openai")
|
|
self.assertEqual(user_config.protocol, "openai_responses")
|
|
self.assertEqual(user_config.base_url, "https://api.openai.com/v1")
|
|
with self.assertRaises(AgentConfigurationError):
|
|
resolve_runtime_model_config(self.bob)
|
|
|
|
def test_disconnect_promotes_remaining_provider(self):
|
|
first = save_verified_provider_config(
|
|
self.alice, self.provider, "safe-test-key", self.result, model_id="model-a"
|
|
)
|
|
second = save_verified_provider_config(
|
|
self.alice,
|
|
get_provider("deepseek"),
|
|
"another-safe-key",
|
|
self.result,
|
|
model_id="model-b",
|
|
)
|
|
self.assertTrue(first.is_default)
|
|
self.assertFalse(second.is_default)
|
|
disconnect_provider(self.alice, "openai")
|
|
second.refresh_from_db()
|
|
self.assertTrue(second.is_default)
|
|
|
|
def test_provider_page_requires_login_and_lists_catalog(self):
|
|
response = self.client.get(reverse("agent_runtime:provider-list"))
|
|
self.assertEqual(response.status_code, 302)
|
|
self.client.force_login(self.alice)
|
|
response = self.client.get(reverse("agent_runtime:provider-list"))
|
|
self.assertContains(response, "阿里云百炼")
|
|
self.assertContains(response, "小米 MiMo")
|
|
self.assertContains(response, "共 23 个内置服务商")
|
|
|
|
@patch("agent_runtime.views.verify_provider_api_key")
|
|
def test_configure_page_saves_verified_key_without_rendering_it(self, verify):
|
|
verify.return_value = self.result
|
|
self.client.force_login(self.alice)
|
|
response = self.client.post(
|
|
reverse("agent_runtime:provider-configure", args=("openai",)),
|
|
{"api_key": "safe-test-key", "model_id": "model-a", "use_as_default": "on"},
|
|
)
|
|
self.assertRedirects(response, reverse("agent_runtime:provider-list"))
|
|
config = ModelProviderConfig.objects.get(owner=self.alice, provider_code="openai")
|
|
self.assertNotIn("safe-test-key", config.encrypted_api_key)
|
|
page = self.client.get(f"{reverse('agent_runtime:provider-list')}?provider=openai")
|
|
self.assertNotContains(page, "safe-test-key")
|
|
self.assertContains(page, config.key_hint)
|
|
|
|
@patch("agent_runtime.views.verify_provider_api_key")
|
|
def test_failed_reconfiguration_does_not_replace_existing_key(self, verify):
|
|
config = save_verified_provider_config(
|
|
self.alice, self.provider, "original-safe-key", self.result, model_id="model-a"
|
|
)
|
|
original_ciphertext = config.encrypted_api_key
|
|
verify.side_effect = ProviderConnectionError("invalid_api_key", "API Key 无效。")
|
|
self.client.force_login(self.alice)
|
|
|
|
response = self.client.post(
|
|
reverse("agent_runtime:provider-configure", args=("openai",)),
|
|
{"api_key": "incorrect-key", "model_id": "model-b"},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
config.refresh_from_db()
|
|
self.assertEqual(config.encrypted_api_key, original_ciphertext)
|
|
self.assertEqual(decrypt_api_key(config.encrypted_api_key), "original-safe-key")
|
|
|
|
def test_connected_filter_never_exposes_another_users_provider(self):
|
|
save_verified_provider_config(
|
|
self.bob, self.provider, "bob-safe-key", self.result, model_id="model-a"
|
|
)
|
|
self.client.force_login(self.alice)
|
|
response = self.client.get(
|
|
reverse("agent_runtime:provider-list"), {"category": "connected"}
|
|
)
|
|
self.assertNotContains(response, "OpenAI")
|