feat: preserve isolated visitor agent runtime
This commit is contained in:
@@ -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):
|
||||||
|
|||||||
@@ -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()))
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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),
|
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
Reference in New Issue
Block a user