"""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, )