feat: preserve isolated visitor agent runtime

This commit is contained in:
张胜宇
2026-09-10 17:36:39 +08:00
parent 8e36a9941d
commit ba22a2220f
10 changed files with 154 additions and 17 deletions
+1
View File
@@ -60,6 +60,7 @@ class AgentDefinition(BaseModel):
allowed_roles: tuple[str, ...] = () allowed_roles: tuple[str, ...] = ()
allowed_portals: tuple[str, ...] = () allowed_portals: tuple[str, ...] = ()
supported_intents: tuple[str, ...] = ("general",) supported_intents: tuple[str, ...] = ("general",)
requires_model_intent_classification: bool = True
class ResolvedAgentConfig(BaseModel): class ResolvedAgentConfig(BaseModel):
+1 -1
View File
@@ -92,6 +92,6 @@ class JwtAuthenticator:
if claims.get("visitor") is True: if claims.get("visitor") is True:
return RequestContext( return RequestContext(
user_id=str(subject), trace_id=str(uuid4()), roles=("visitor",), user_id=str(subject), trace_id=str(uuid4()), roles=("visitor",),
permissions=("agent:run",), permissions=("agent:run",), data_scope="public",
) )
return RequestContext(user_id=str(claims["sub"]), trace_id=str(uuid4())) return RequestContext(user_id=str(claims["sub"]), trace_id=str(uuid4()))
+5
View File
@@ -129,11 +129,16 @@ class BaseAgent(ABC):
async def recall_memory(self, request: AgentRequest, context: RequestContext) -> None: async def recall_memory(self, request: AgentRequest, context: RequestContext) -> None:
if self._governance is None: if self._governance is None:
raise RecoverableAgentError("缺少记忆治理依赖") raise RecoverableAgentError("缺少记忆治理依赖")
if "visitor" in context.roles:
self.memories = ()
return
self.memories = await self._governance.recall(context) self.memories = await self._governance.recall(context)
if any(memory.customer_id != context.user_id for memory in self.memories): if any(memory.customer_id != context.user_id for memory in self.memories):
raise RecoverableAgentError("记忆召回越过客户范围") raise RecoverableAgentError("记忆召回越过客户范围")
async def classify_intent(self, request: AgentRequest) -> IntentResult | None: async def classify_intent(self, request: AgentRequest) -> IntentResult | None:
if not self.definition.requires_model_intent_classification:
return None
if self._intent_classifier is None or self._intent_endpoint_resolver is None: if self._intent_classifier is None or self._intent_endpoint_resolver is None:
return None return None
endpoints = await self._intent_endpoint_resolver.resolve( endpoints = await self._intent_endpoint_resolver.resolve(
+5 -1
View File
@@ -59,6 +59,10 @@ class AgentFactory:
agent.bind_model_service(self._model_service) agent.bind_model_service(self._model_service)
if self._tool_executor is not None: if self._tool_executor is not None:
agent.bind_tool_executor(self._tool_executor) agent.bind_tool_executor(self._tool_executor)
if self._intent_classifier is not None and self._intent_endpoint_resolver is not None: if (
agent.definition.requires_model_intent_classification
and self._intent_classifier is not None
and self._intent_endpoint_resolver is not None
):
agent.bind_intent_classifier(self._intent_classifier, self._intent_endpoint_resolver) agent.bind_intent_classifier(self._intent_classifier, self._intent_endpoint_resolver)
return agent return agent
+5 -1
View File
@@ -118,7 +118,11 @@ class AgentRunApplicationService:
await OutboxRepository(self.session).append(DomainEvent( await OutboxRepository(self.session).append(DomainEvent(
event_id=str(uuid4()), event_type="agent.run_requested", aggregate_type="agent_run", event_id=str(uuid4()), event_type="agent.run_requested", aggregate_type="agent_run",
aggregate_id=run_id, trace_id=trace_id, aggregate_id=run_id, trace_id=trace_id,
payload={"run_id": run_id, "metadata": request.metadata.model_dump(mode="json")}, payload={
"run_id": run_id,
"actor_type": "visitor" if "visitor" in context.roles else "authenticated",
"metadata": request.metadata.model_dump(mode="json"),
},
occurred_at=now, occurred_at=now,
)) ))
return RunAccepted(run_id, trace_id) return RunAccepted(run_id, trace_id)
+43 -14
View File
@@ -102,6 +102,35 @@ class WorkerRuntime:
# episode 聚合是低频批处理,按轮次节流而不是每轮都查。 # episode 聚合是低频批处理,按轮次节流而不是每轮都查。
self._episode_rounds = 0 self._episode_rounds = 0
async def restore_context(
self, *, actor_type: str, actor_id: str, trace_id: str
) -> RequestContext:
"""按已验证的内部事件身份恢复执行上下文。"""
identity = RequestContext(user_id=actor_id, trace_id=trace_id)
if actor_type == "visitor":
return identity.model_copy(update={
"roles": ("visitor",),
"permissions": ("agent:run",),
"data_scope": "public",
})
return await self.resolve_identity(identity)
@staticmethod
def should_request_memory_extraction(
*, context: RequestContext, message: str, result: AgentResult,
business_events: tuple[str, ...] | list[str],
) -> bool:
"""只允许已登录用户的明确业务事实进入客户记忆抽取队列。"""
if "visitor" in context.roles:
return False
return MemoryService.should_extract_memory(
conversation_content=message,
role="user",
tool_result=any(call.status == "succeeded" for call in result.result.tool_calls),
event_type=business_events[0] if business_events else None,
signals=MemoryService.detect_memory_signals(message),
)
async def dispatch_one(self, *, run_id: str | None = None) -> bool: async def dispatch_one(self, *, run_id: str | None = None) -> bool:
# Outbox acknowledges a durable SQL queue entry, not an in-memory task. # Outbox acknowledges a durable SQL queue entry, not an in-memory task.
async with SessionFactory() as session: async with SessionFactory() as session:
@@ -428,9 +457,15 @@ class WorkerRuntime:
idempotency_key=idem.idempotency_key, idempotency_key=idem.idempotency_key,
metadata=AgentRequestMetadata.model_validate(metadata), metadata=AgentRequestMetadata.model_validate(metadata),
) )
identity = RequestContext(user_id=str(run.user_id), trace_id=run.trace_id) actor_type = (
# Re-check account and permissions at execution time, including delayed jobs. str(event.payload.get("actor_type", "authenticated"))
context = await self.resolve_identity(identity) if event else "authenticated"
)
actor_id = str(run.user_id)
trace_id = run.trace_id
context = await self.restore_context(
actor_type=actor_type, actor_id=actor_id, trace_id=trace_id
)
result: AgentResult | None = None result: AgentResult | None = None
async for event_data in AgentExecutor(self.factory).execute( async for event_data in AgentExecutor(self.factory).execute(
request.agent_type, request, context, run_id request.agent_type, request, context, run_id
@@ -452,17 +487,11 @@ class WorkerRuntime:
async with SessionFactory() as session: async with SessionFactory() as session:
await AgentPersistenceService(session).complete_run( await AgentPersistenceService(session).complete_run(
run_id, result, worker_id=worker_id, run_id, result, worker_id=worker_id,
memory_extraction_requested=MemoryService.should_extract_memory( memory_extraction_requested=self.should_request_memory_extraction(
conversation_content=request.message, context=context,
role="user", message=request.message,
# 工具产出的权威事实同样构成持久记忆(工具调用记录来自终态结果)。 result=result,
tool_result=any( business_events=business_events,
call.status == "succeeded" for call in result.result.tool_calls
),
# 本 run 落库的业务事件(风险评估完成、交易完成等)。
event_type=business_events[0] if business_events else None,
# 用户明确陈述的偏好/约束/身份/目标,命中才触发抽取。
signals=MemoryService.detect_memory_signals(request.message),
), ),
) )
+1
View File
@@ -55,6 +55,7 @@ def test_authenticate_visitor_token_returns_limited_anonymous_context() -> None:
assert context.roles == ("visitor",) assert context.roles == ("visitor",)
assert context.permissions == ("agent:run",) assert context.permissions == ("agent:run",)
assert context.customer_ids == () assert context.customer_ids == ()
assert context.data_scope == "public"
def test_visitor_token_issuer_creates_short_lived_limited_token() -> None: def test_visitor_token_issuer_creates_short_lived_limited_token() -> None:
@@ -23,6 +23,36 @@ def test_all_governance_hooks_protected(name):
type("Bypass", (BaseAgent,), {name: lambda *args: None}) type("Bypass", (BaseAgent,), {name: lambda *args: None})
@pytest.mark.asyncio
async def test_visitor_does_not_recall_customer_memory() -> None:
"""访客不能以匿名主体标识读取任何客户记忆。"""
class Demo(BaseAgent):
async def handle(self, request, context):
return CoreResult(text="unused")
class FailingGovernance:
async def recall(self, context):
raise AssertionError("visitor memory recall is forbidden")
definition = AgentDefinition(
agent_type="demo", version="1", allowed_roles=("visitor",), allowed_portals=("api",)
)
agent = Demo(definition)
agent.bind_governance(FailingGovernance())
request = AgentRequest(
agent_type="demo", message="公开问题", session_id="visitor-session",
idempotency_key="visitor-memory-request-0001",
)
context = RequestContext(
user_id="visitor-id", trace_id="visitor-trace", roles=("visitor",),
permissions=("agent:run",), data_scope="public",
)
await agent.recall_memory(request, context)
assert agent.memories == ()
async def test_resolve_recall_handle_review_order_and_snapshot(governance): async def test_resolve_recall_handle_review_order_and_snapshot(governance):
calls = [] calls = []
config = ResolvedAgentConfig(config_version="released", prompt_version="p", model_endpoint="m") config = ResolvedAgentConfig(config_version="released", prompt_version="p", model_endpoint="m")
@@ -51,6 +51,34 @@ async def test_execute_classifies_before_handle_and_attaches_result(governance)
assert result["intent"]["confidence"] == 0.9 assert result["intent"]["confidence"] == 0.9
@pytest.mark.asyncio
async def test_fixed_route_agent_skips_model_intent_classification(governance) -> None:
"""固定路由 Agent 不能因模型意图端点不可用而阻断服务。"""
definition = AgentDefinition(
agent_type="demo", version="1", allowed_roles=("customer",),
allowed_portals=("api",), requires_model_intent_classification=False,
)
factory = AgentFactory(
governance=governance,
intent_classifier=IntentClassifier(StubModel()),
intent_endpoint_resolver=StubResolver(),
)
factory.register(definition, lambda _context: DemoAgent(definition))
context = RequestContext(
user_id="1", trace_id="fixed-route", roles=("customer",), permissions=("agent:run",)
)
request = AgentRequest(
agent_type="demo", message="本地路由", session_id="s",
idempotency_key="fixed-route-request-0001",
)
events = [
event async for event in factory.create("demo", context).execute(request, context, "run")
]
assert events[-1].payload["result"]["result"]["intent"] is None
def test_business_agent_cannot_override_intent_governance() -> None: def test_business_agent_cannot_override_intent_governance() -> None:
with pytest.raises(TypeError, match="classify_intent"): with pytest.raises(TypeError, match="classify_intent"):
class InvalidAgent(BaseAgent): class InvalidAgent(BaseAgent):
@@ -4,6 +4,7 @@ from unittest.mock import AsyncMock
import pytest import pytest
from app.core.contracts import AgentResult, CoreResult, RequestContext
from app.service.memory_recall_service import MemoryRecallService from app.service.memory_recall_service import MemoryRecallService
from app.service.model_gateway import ModelGenerationService from app.service.model_gateway import ModelGenerationService
from app.worker.runtime import WorkerRuntime from app.worker.runtime import WorkerRuntime
@@ -22,6 +23,40 @@ OUTBOX = {
} }
@pytest.mark.asyncio
async def test_worker_restores_visitor_without_identity_repository_call() -> None:
"""访客异步任务只能恢复最小公开上下文,不能查询正式身份库。"""
runtime = WorkerRuntime.__new__(WorkerRuntime)
runtime.resolve_identity = AsyncMock(side_effect=AssertionError("identity lookup is forbidden"))
context = await runtime.restore_context(
actor_type="visitor", actor_id="visitor:test", trace_id="trace-test"
)
assert context.roles == ("visitor",)
assert context.permissions == ("agent:run",)
assert context.data_scope == "public"
runtime.resolve_identity.assert_not_awaited()
def test_visitor_does_not_request_memory_extraction() -> None:
"""访客消息即使包含偏好信号,也不能进入客户记忆抽取队列。"""
context = RequestContext(
user_id="visitor:test", trace_id="visitor-trace", roles=("visitor",),
permissions=("agent:run",), data_scope="public",
)
result = AgentResult(run_id="visitor-run", result=CoreResult(text="公开答复"))
requested = WorkerRuntime.should_request_memory_extraction(
context=context,
message="我的风险偏好是稳健型",
result=result,
business_events=(),
)
assert requested is False
class FakeSession(AbstractAsyncContextManager["FakeSession"]): class FakeSession(AbstractAsyncContextManager["FakeSession"]):
def __init__(self) -> None: def __init__(self) -> None:
self.scalar = AsyncMock(return_value="event-1") self.scalar = AsyncMock(return_value="event-1")