184 lines
7.8 KiB
Python
184 lines
7.8 KiB
Python
"""RC-BE-007: DeerFlow model freeze + AgentScope chat adapter."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from types import SimpleNamespace
|
|
|
|
import pytest
|
|
|
|
from app.report_collaboration.agentscope_runtime.event_adapter import ConversationReplica, MemberStreamContext, SeqAllocator, make_emitter
|
|
from app.report_collaboration.agentscope_runtime.model_adapter import DeerFlowChatModelAdapter, ModelAdapterError, record_attempt_usage, redact_secrets
|
|
from app.report_collaboration.agentscope_runtime.model_resolver import freeze_plan_models, resolve_role_model
|
|
from deerflow.config.app_config import AppConfig
|
|
from deerflow.config.model_config import ModelConfig
|
|
from deerflow.config.report_collaboration_config import ReportCollaborationConfig
|
|
from deerflow.persistence.report_collaboration import ReportCollaborationValidationError
|
|
|
|
|
|
def _app(*, models: list[ModelConfig] | None = None, rc: ReportCollaborationConfig | None = None) -> AppConfig:
|
|
return AppConfig.model_validate(
|
|
{
|
|
"sandbox": {"use": "deerflow.sandbox.local:LocalSandboxProvider"},
|
|
"models": [item.model_dump() for item in (models or [ModelConfig(name="demo-chat", use="langchain_openai:ChatOpenAI", model="demo")])],
|
|
"report_collaboration": (rc or ReportCollaborationConfig()).model_dump(),
|
|
}
|
|
)
|
|
|
|
|
|
def _plan(*, mission: str = "检索证据", role: str = "researcher") -> dict:
|
|
return {
|
|
"nodes": [
|
|
{
|
|
"id": "research-0",
|
|
"label": "研究",
|
|
"role_key": role,
|
|
"mission": mission,
|
|
"output_artifact_type": "EvidenceBundle",
|
|
"allowed_tools": ["web_search"],
|
|
"acceptance_criteria": ["有来源"],
|
|
"depends_on": [],
|
|
}
|
|
],
|
|
"roles": [{"key": role, "display_name": "研究员", "responsibility": "检索"}],
|
|
}
|
|
|
|
|
|
class _FakeLLM:
|
|
def __init__(self, chunks: list | None = None, error: Exception | None = None) -> None:
|
|
self.chunks = chunks or []
|
|
self.error = error
|
|
self.bound_tools = None
|
|
|
|
def bind_tools(self, tools):
|
|
self.bound_tools = tools
|
|
return self
|
|
|
|
async def astream(self, messages):
|
|
_ = messages
|
|
if self.error is not None:
|
|
raise self.error
|
|
for chunk in self.chunks:
|
|
yield chunk
|
|
|
|
|
|
def test_named_model_missing_fails_before_run() -> None:
|
|
app = _app(rc=ReportCollaborationConfig(research_model="missing-model"))
|
|
with pytest.raises(ReportCollaborationValidationError) as exc:
|
|
freeze_plan_models(_plan(), app)
|
|
assert exc.value.code == "MODEL_NOT_FOUND"
|
|
|
|
|
|
def test_unspecified_role_falls_back_and_is_recorded() -> None:
|
|
planner = ModelConfig(name="planner", use="langchain_openai:ChatOpenAI", model="p1")
|
|
app = _app(models=[planner], rc=ReportCollaborationConfig(planner_model="planner"))
|
|
snapshot = resolve_role_model("researcher", app_config=app)
|
|
assert snapshot.config_name == "planner"
|
|
assert snapshot.fallback_reason == "role_unspecified_use_planner"
|
|
assert snapshot.frozen is True
|
|
dumped = snapshot.model_dump()
|
|
assert "api_key" not in dumped
|
|
assert "base_url" not in dumped
|
|
|
|
|
|
def test_vision_capability_mismatch() -> None:
|
|
app = _app()
|
|
with pytest.raises(ReportCollaborationValidationError) as exc:
|
|
freeze_plan_models(_plan(mission="分析配图与图像来源"), app)
|
|
assert exc.value.code == "MODEL_CAPABILITY_MISMATCH"
|
|
|
|
|
|
def test_snapshot_strips_provider_secrets() -> None:
|
|
model = ModelConfig(name="secret-chat", use="langchain_openai:ChatOpenAI", model="gpt", api_key="sk-SECRETVALUE", base_url="https://hidden.example")
|
|
app = _app(models=[model], rc=ReportCollaborationConfig(research_model="secret-chat"))
|
|
bundle = freeze_plan_models(_plan(), app)
|
|
blob = str(bundle.model_dump())
|
|
assert "sk-SECRETVALUE" not in blob
|
|
assert "hidden.example" not in blob
|
|
assert bundle.by_role["researcher"].config_name == "secret-chat"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_adapter_streams_text_tool_and_usage() -> None:
|
|
snapshot = resolve_role_model("researcher", app_config=_app())
|
|
llm = _FakeLLM(
|
|
chunks=[
|
|
SimpleNamespace(content="市", tool_call_chunks=[], usage_metadata=None, response_metadata={"api_key": "sk-LEAK"}),
|
|
SimpleNamespace(
|
|
content="场",
|
|
tool_call_chunks=[SimpleNamespace(id="call_1", name="web_search", args='{"q":')],
|
|
usage_metadata={"input_tokens": 12, "output_tokens": 4, "total_tokens": 16},
|
|
response_metadata={},
|
|
),
|
|
]
|
|
)
|
|
adapter = DeerFlowChatModelAdapter(snapshot, app_config=_app(), llm_factory=lambda: llm, node_run_id="research-0-attempt1")
|
|
pieces = [item async for item in adapter.astream_chat([{"role": "user", "content": "写报告"}], tools=[{"name": "web_search"}])]
|
|
assert llm.bound_tools == [{"name": "web_search"}]
|
|
assert any(not item.is_last and item.content and item.content[0].get("text") == "市" for item in pieces)
|
|
last = pieces[-1]
|
|
assert last.is_last is True
|
|
assert last.usage is not None
|
|
assert last.usage.input_tokens == 12
|
|
assert adapter.last_usage.total_tokens == 16
|
|
assert adapter.last_usage.node_run_id == "research-0-attempt1"
|
|
assert "[REDACTED]" in str(pieces[0].metadata.get("api_key"))
|
|
adapter.refuse_swap("demo-chat")
|
|
with pytest.raises(ModelAdapterError) as frozen:
|
|
adapter.refuse_swap("other-model")
|
|
assert frozen.value.code == "MODEL_FROZEN"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_adapter_redacts_provider_errors_and_records_metrics() -> None:
|
|
snapshot = resolve_role_model("writer", app_config=_app())
|
|
llm = _FakeLLM(error=RuntimeError("invalid api_key sk-SECRETVALUE Bearer abc.def"))
|
|
adapter = DeerFlowChatModelAdapter(snapshot, app_config=_app(), llm_factory=lambda: llm)
|
|
with pytest.raises(ModelAdapterError) as exc:
|
|
_ = [item async for item in adapter.astream_chat([{"role": "user", "content": "hi"}])]
|
|
assert "sk-SECRETVALUE" not in str(exc.value)
|
|
assert "[REDACTED]" in str(exc.value)
|
|
assert adapter.last_usage.status == "error"
|
|
assert "sk-SECRETVALUE" not in (adapter.last_usage.error_message or "")
|
|
|
|
class _Metrics:
|
|
def __init__(self) -> None:
|
|
self.rows: list = []
|
|
|
|
async def record(self, metric) -> None:
|
|
self.rows.append(metric)
|
|
|
|
metrics = _Metrics()
|
|
await record_attempt_usage(adapter.last_usage, user_id="u1", run_id="run_1", metrics_store=metrics)
|
|
assert metrics.rows[0].status == "error"
|
|
assert "sk-SECRETVALUE" not in (metrics.rows[0].error_message or "")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_protocol_events_feed_conversation_replica() -> None:
|
|
snapshot = resolve_role_model("researcher", app_config=_app())
|
|
llm = _FakeLLM(chunks=[SimpleNamespace(content="份额上升", tool_call_chunks=[], usage_metadata={"input_tokens": 1, "output_tokens": 2}, response_metadata={})])
|
|
adapter = DeerFlowChatModelAdapter(snapshot, app_config=_app(), llm_factory=lambda: llm)
|
|
replica = ConversationReplica(
|
|
MemberStreamContext(
|
|
session_id="s1",
|
|
run_id="r1",
|
|
agent_run_id="agent_1",
|
|
node_run_id="research-0-attempt1",
|
|
phase_id="p1",
|
|
display_name="研究员",
|
|
),
|
|
make_emitter("s1", "r1", SeqAllocator()),
|
|
)
|
|
types: list[str] = []
|
|
async for event in adapter.astream_protocol_events([{"role": "user", "content": "开始"}], display_name="研究员"):
|
|
emitted = replica.ingest(event)
|
|
types.extend(item.type for item in emitted)
|
|
assert types[0] == "message.created"
|
|
assert "message.delta" in types
|
|
assert types[-1] == "message.completed"
|
|
|
|
|
|
def test_redact_helper() -> None:
|
|
assert "sk-abc" not in redact_secrets("using sk-abc and Bearer tok")
|
|
assert "[REDACTED]" in redact_secrets("api_key=super-secret")
|