289 lines
8.9 KiB
Python
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,
|
|
)
|