deerflow-code/offline-backend-20260512/backend/tests/test_fixed_questions.py
2026-09-07 18:24:55 +08:00

289 lines
8.9 KiB
Python

"""Regression tests for administrator-managed exact-match fixed answers."""
from __future__ import annotations
from types import SimpleNamespace
import pytest
from langchain_core.messages import HumanMessage
from pydantic import ValidationError
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
from app.gateway.routers.fixed_questions import (
FixedQuestion,
FixedQuestionReplaceRequest,
)
from app.gateway.services import (
RUN_SCOPE_EXTERNAL_ROOT,
RUN_SCOPE_ORCHESTRATION_CHILD,
_find_fixed_answer,
normalize_input,
)
from deerflow.agents import fixed_answer_agent as fixed_answer_module
from deerflow.agents.fixed_answer_agent import make_fixed_answer_agent
from deerflow.persistence.fixed_questions import MemoryFixedQuestionStore
from deerflow.persistence.fixed_questions.model import FixedQuestionRow
from deerflow.persistence.fixed_questions.sql import FixedQuestionRepository
def _body(**overrides):
data = {
"assistant_id": "lead_agent",
"command": None,
"context": {},
}
data.update(overrides)
return SimpleNamespace(**data)
def _request(store: MemoryFixedQuestionStore):
return SimpleNamespace(
app=SimpleNamespace(
state=SimpleNamespace(fixed_question_store=store),
)
)
@pytest.mark.asyncio
async def test_memory_store_matches_only_enabled_exact_question() -> None:
store = MemoryFixedQuestionStore()
await store.replace_all(
[
{
"id": "enabled",
"question": "系统开放时间?",
"answer": "工作日 09:00—18:00。",
"enabled": True,
"tokens_per_second": 80,
"sort_order": 0,
},
{
"id": "disabled",
"question": "停用问题",
"answer": "不应返回",
"enabled": False,
"sort_order": 1,
},
]
)
match = await store.find_enabled("系统开放时间?")
assert match is not None
assert match["answer"] == "工作日 09:00—18:00。"
assert match["tokens_per_second"] == 80
assert await store.find_enabled("系统开放时间?") is None
assert await store.find_enabled(" 系统开放时间?") is None
assert await store.find_enabled("停用问题") is None
@pytest.mark.asyncio
async def test_sql_store_round_trips_and_exact_matches() -> None:
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
try:
async with engine.begin() as connection:
await connection.run_sync(FixedQuestionRow.__table__.create)
repository = FixedQuestionRepository(async_sessionmaker(engine, expire_on_commit=False))
saved = await repository.replace_all(
[
{
"id": "sql-q1",
"question": "SQL 精确问题",
"answer": "SQL 固定回答",
"enabled": True,
"tokens_per_second": 120,
"sort_order": 0,
}
],
updated_by="admin",
)
assert saved[0]["updated_by"] == "admin"
assert saved[0]["tokens_per_second"] == 120
assert (await repository.find_enabled("SQL 精确问题"))["answer"] == "SQL 固定回答"
assert await repository.find_enabled("sql 精确问题") is None
finally:
await engine.dispose()
@pytest.mark.asyncio
async def test_lookup_skips_attachments_and_internal_runs() -> None:
store = MemoryFixedQuestionStore()
await store.replace_all(
[
{
"id": "q1",
"question": "固定问题",
"answer": "固定回答",
"enabled": True,
"sort_order": 0,
}
]
)
request = _request(store)
plain_input = normalize_input({"messages": [{"content": "固定问题"}]})
match = await _find_fixed_answer(
body=_body(),
request=request,
graph_input=plain_input,
run_scope=RUN_SCOPE_EXTERNAL_ROOT,
)
assert match is not None
assert match["id"] == "q1"
attachment_input = normalize_input(
{
"messages": [
{
"content": "固定问题",
"additional_kwargs": {"files": [{"name": "context.pdf"}]},
}
]
}
)
assert (
await _find_fixed_answer(
body=_body(),
request=request,
graph_input=attachment_input,
run_scope=RUN_SCOPE_EXTERNAL_ROOT,
)
is None
)
contextual_input = normalize_input(
{
"messages": [
{
"content": "固定问题",
"additional_kwargs": {"prompt_prefix": "网页正文上下文"},
}
]
}
)
assert (
await _find_fixed_answer(
body=_body(),
request=request,
graph_input=contextual_input,
run_scope=RUN_SCOPE_EXTERNAL_ROOT,
)
is None
)
assert (
await _find_fixed_answer(
body=_body(),
request=request,
graph_input=plain_input,
run_scope=RUN_SCOPE_ORCHESTRATION_CHILD,
)
is None
)
@pytest.mark.asyncio
async def test_fixed_answer_agent_streams_chunks_and_persists_full_answer() -> None:
question = "固定问题"
answer = "这是一个足够长的固定回答,用于验证返回内容会拆成多个流式消息块,同时完整写入会话状态。"
agent = make_fixed_answer_agent(
question=question,
answer=answer,
fixed_question_id="q1",
tokens_per_second=1000,
)
chunks: list[str] = []
final_state = None
async for mode, data in agent.astream(
{"messages": [HumanMessage(content=question)]},
stream_mode=["values", "messages"],
):
if mode == "messages":
message, _metadata = data
if isinstance(message.content, str):
chunks.append(message.content)
elif mode == "values":
final_state = data
assert "".join(chunks) == answer
assert len([chunk for chunk in chunks if chunk]) >= 2
assert final_state is not None
assert final_state["messages"][-1].content == answer
assert final_state["messages"][-1].additional_kwargs["fixed_answer"] is True
assert final_state["messages"][-1].additional_kwargs["tokens_per_second"] == 1000
assert final_state["title"] == question
assert final_state["title_provisional"] is False
@pytest.mark.asyncio
async def test_fixed_answer_default_speed_schedules_100_tokens_per_second(
monkeypatch: pytest.MonkeyPatch,
) -> None:
answer = "one two three four five six seven eight nine ten eleven twelve thirteen"
sleeps: list[float] = []
async def record_sleep(delay: float) -> None:
sleeps.append(delay)
monkeypatch.setattr(fixed_answer_module.asyncio, "sleep", record_sleep)
model = fixed_answer_module.FixedAnswerChatModel(
answer=answer,
message_id="speed-test",
fixed_question_id="q-speed",
)
chunks: list[str] = []
async for generation in model._astream([]):
chunks.append(str(generation.message.content))
token_count = len(fixed_answer_module._TOKEN_ENCODING.encode(answer))
assert "".join(chunks) == answer
assert model.tokens_per_second == 100
assert sum(sleeps) == pytest.approx((token_count - 1) / 100)
def test_replace_request_rejects_duplicate_exact_questions() -> None:
row = {
"question": "重复问题",
"answer": "回答",
"enabled": True,
"tokens_per_second": 100,
"sort_order": 0,
}
with pytest.raises(ValidationError):
FixedQuestionReplaceRequest(
questions=[
FixedQuestion(id="q1", **row),
FixedQuestion(id="q2", **row),
]
)
with pytest.raises(ValidationError):
FixedQuestionReplaceRequest(
questions=[
FixedQuestion(id="same-id", **row),
FixedQuestion(
id="same-id",
question="另一个问题",
answer="另一个回答",
enabled=True,
sort_order=1,
),
]
)
def test_fixed_question_stream_speed_defaults_and_bounds() -> None:
default_row = FixedQuestion(
id="default-speed",
question="默认速度",
answer="回答",
)
assert default_row.tokens_per_second == 100
for invalid_speed in (0, 1001):
with pytest.raises(ValidationError):
FixedQuestion(
id=f"invalid-{invalid_speed}",
question="非法速度",
answer="回答",
tokens_per_second=invalid_speed,
)