"""Approved cache: no external Redis, LLM or Neo4j connections.""" from copy import deepcopy import json from types import SimpleNamespace import unittest from unittest.mock import patch from fastapi import HTTPException from langgraph.checkpoint.memory import InMemorySaver from step3_qa_agent.agent import nodes, production from step3_qa_agent.agent.graph import build_agent_graph from step4_web import api from step4_web.answer_cache import (ApprovedAnswerCache, prepare_request, cache_key, eligibility, restored_state, review_info, ACTIVE_VERSION, APPROVE_CURRENT) from scripts.test_simple_plan import CONTEXT, state as simple_state from step3_qa_agent.agent.simple_plan import build_simple_plan RELEASE = {'data_version': 'v1', 'schema_version': 'schema1'} Q = '共有多少名员工?' class FakeRedis: def __init__(self): self.entries = {ACTIVE_VERSION: 'v1'} self.ttls = {} self.writes = 0 async def get(self, key): return self.entries.get(key) async def set(self, key, value, ex=None): self.entries[key] = value self.ttls[key] = ex self.writes += 1 async def eval(self, script, count, key, *args): if script == APPROVE_CURRENT: fence, raw, version = args if self.entries.get(fence) != version: return 0 await self.set(key, raw) return 1 raw, = args if self.entries.get(key) == raw: del self.entries[key] return 1 return 0 def completed(): s = simple_state() s.update(answer='共有2名员工,按工号去重。', answer_id='answer1', fit=True, cache_request=prepare_request('t1', Q, {}, RELEASE), subgraph={'results': {'s1': {'rows': [{'记录数': 2}], 'truncated': False, 'aggregate': True, 'data_version': 'v1'}}, 'suggestions': {'options': []}}, trace=[{'node': 'understand'}]) s['plan'] = build_simple_plan(s)[0] return s class CacheTests(unittest.IsolatedAsyncioTestCase): async def asyncSetUp(self): self.redis = FakeRedis() self.cache = ApprovedAnswerCache(self.redis) async def test_lookup_never_populates_cache(self): self.assertIsNone(await self.cache.lookup(completed()['cache_request'])) self.assertEqual(self.redis.writes, 0) async def test_only_approved_answer_hits_permanently(self): s = completed() entry = await self.cache.approve(s) loaded = await self.cache.lookup(s['cache_request']) self.assertEqual(loaded['payload']['answer'], s['answer']) self.assertTrue(loaded['approved']) self.assertIsNone(self.redis.ttls[cache_key(s['cache_request'])]) self.assertNotIn('expires_at', entry) self.assertTrue(cache_key(s['cache_request']).startswith('ka:qa:approved:v1:')) self.assertEqual(entry['entry_id'], 'answer1') async def test_old_inflight_approval_rejected_after_version_switch(self): self.redis.entries[ACTIVE_VERSION] = 'v2' with self.assertRaises(ValueError): await self.cache.approve(completed()) self.assertEqual(self.redis.writes, 0) async def test_scope_version_schema_context_and_question_are_separate(self): s = completed() await self.cache.approve(s) for field in ('scope', 'data_version', 'schema_version', 'context', 'question'): request = {**s['cache_request'], field: 'different'} self.assertIsNone(await self.cache.lookup(request)) async def test_no_semantic_or_negation_matching(self): s = completed() await self.cache.approve(s) for q in ('员工共有几人?', '没有证书的员工有多少?', '共有多少名员工!'): self.assertIsNone(await self.cache.lookup(prepare_request('t1', q, {}, RELEASE))) async def test_contextual_and_relative_questions_are_not_cacheable(self): for question in ('他们有多少人', '这些员工的岗位', '今年入职人数'): s = completed() s['question'] = question s['cache_request']['question'] = question self.assertTrue(eligibility(s)) with self.assertRaises(ValueError): await self.cache.approve(s) self.assertEqual(self.redis.writes, 0) async def test_failed_truncated_modified_and_partial_answers_rejected(self): variations = [] s = completed(); s['fit'] = False; variations.append(s) s = completed(); s['category'] = '闲聊'; variations.append(s) s = completed(); s['plan_feedback'] = '换成项目'; variations.append(s) s = completed(); s['subgraph']['results']['s1']['truncated'] = True; variations.append(s) s = completed(); s['subgraph']['results']['s1']['data_version'] = 'old'; variations.append(s) s = completed(); s['subgraph']['results'] = {}; variations.append(s) s = completed(); s['subgraph']['error'] = 'failed'; variations.append(s) s = completed(); s['subgraph']['suggestions']['options'] = ['candidate']; variations.append(s) s = completed(); s['trace'].append({'node': 'value_choice'}); variations.append(s) s = completed(); s['answer'] = ''; variations.append(s) for s in variations: self.assertTrue(eligibility(s)) with self.assertRaises(ValueError): await self.cache.approve(s) self.assertEqual(self.redis.writes, 0) async def test_corrupt_unapproved_entries_miss(self): s = completed() entry = await self.cache.approve(s) key = cache_key(s['cache_request']) for value in ('broken', '[]', json.dumps({**entry, 'approved': False})): self.redis.entries[key] = value self.assertIsNone(await self.cache.lookup(s['cache_request'])) async def test_cache_read_timeout_falls_back(self): import asyncio async def slow(key): await asyncio.sleep(1) self.cache.timeout_seconds = .001 with patch.object(self.redis, 'get', side_effect=slow): self.assertIsNone(await self.cache.lookup(completed()['cache_request'])) async def test_revoke_and_old_revoke_cannot_delete_new_approval(self): s = completed() await self.cache.approve(s) s['answer_id'] = 'new-answer' await self.cache.approve(s) with self.assertRaises(ValueError): await self.cache.revoke(s['cache_request'], 'answer1') self.assertIsNotNone(await self.cache.lookup(s['cache_request'])) await self.cache.revoke(s['cache_request'], 'new-answer') self.assertIsNone(await self.cache.lookup(s['cache_request'])) async def test_repeat_preserves_original_context_but_other_context_misses(self): first = prepare_request('t1', Q, {}, RELEASE) prev = {'rounds': [{'question': Q, 'answer': '2', 'schema_build_id': 'v1', 'cache_context': first['context']}], 'messages': [{'role': 'user', 'content': Q}, {'role': 'assistant', 'content': '2'}]} self.assertEqual(prepare_request('t1', Q, prev, RELEASE), first) prev['rounds'].append({'question': '只看保安', 'answer': '1'}) self.assertNotEqual(prepare_request('t1', Q, prev, RELEASE), first) async def test_oversize_not_eligible_and_never_written(self): s = completed(); s['answer'] = 'x' * 300000 self.assertFalse(review_info(s, 'cp')['eligible']) with self.assertRaises(ValueError): await self.cache.approve(s) self.assertEqual(self.redis.writes, 0) async def test_neo4j_records_keep_named_columns(self): from neo4j import Record s = completed() s['subgraph']['results']['s1']['rows'] = [Record({'记录数': 2})] await self.cache.approve(s) hit = await self.cache.lookup(s['cache_request']) self.assertEqual(hit['payload']['subgraph']['results']['s1']['rows'], [{'记录数': 2}]) class ApiCacheTests(unittest.IsolatedAsyncioTestCase): async def asyncSetUp(self): self.redis = FakeRedis() self.cache = ApprovedAnswerCache(self.redis) s = simple_state() self.planner = patch.object(production, 'chat_json', return_value=build_simple_plan(s)[0]) self.patches = [ patch.object(api, '_graph', build_agent_graph(InMemorySaver())), patch.object(api, '_answer_cache', self.cache), patch.object(api, 'release_snapshot', return_value=RELEASE), patch.object(api, '_write_api_log'), patch.object(nodes, 'load_question_context', return_value=deepcopy(CONTEXT)), patch.object(nodes, 'select_assessments', side_effect=lambda ctx, *a, **kw: ctx), patch.object(nodes, 'chat_json', return_value={'category': '图谱检索', 'slots': {}, 'data_needs': s['simple_query_needs']}), self.planner, patch.object(production, 'chat_text', return_value='共有2名员工,按工号去重。'), patch.object(production, 'release_snapshot', return_value={'data_version': 'v1'}), patch.object(production, 'get_driver', return_value=SimpleNamespace( execute_query=lambda *a, **kw: SimpleNamespace(records=[{'记录数': 2}])))] self.mocks = [p.start() for p in self.patches] self.addCleanup(lambda: [p.stop() for p in reversed(self.patches)]) async def ask(self, question=Q, reuse=True, auto=True): events = [e async for e in api._astream_run('t1', question, auto, reuse)] return events[-1] def body(self, state, **extra): review = state['cache_review'] return api.AnswerReviewRequest(answer_id=review['answer_id'], checkpoint_id=review['checkpoint_id'], **extra) async def test_cache_miss_reports_understanding_before_model_starts(self): stream = api._astream_run('t1', Q, True, True) try: self.assertEqual((await anext(stream))['node'], 'answer_cache') self.assertEqual((await anext(stream))['node'], 'understand') self.assertEqual(self.mocks[6].call_count, 0) self.assertEqual(self.mocks[2].call_count, 1) remaining = [event async for event in stream] self.assertEqual(remaining[-1]['type'], 'done') finally: await stream.aclose() async def test_full_flow_requires_click_then_repeat_uses_no_models(self): first = (await self.ask())['state'] self.assertTrue(first['cache_review']['eligible']) self.assertEqual(self.redis.writes, 0) result = await api.approve_answer('t1', self.body(first)) self.assertEqual(self.redis.writes, 1) calls = [self.mocks[i].call_count for i in (6, 7, 8)] hit = (await self.ask())['state'] self.assertTrue(hit['cache_review']['hit']) self.assertEqual(hit['answer'], first['answer']) self.assertEqual([self.mocks[i].call_count for i in (6, 7, 8)], calls) self.assertEqual(len(hit['rounds']), 2) self.assertNotEqual(hit['answer_id'], first['answer_id']) await api.revoke_answer('t1', self.body(hit, entry_id=result['entry_id'])) fresh = (await self.ask())['state'] self.assertFalse(fresh['cache_review']['hit']) async def test_reuse_switch_bypasses_cache(self): first = (await self.ask())['state'] await api.approve_answer('t1', self.body(first)) next_answer = (await self.ask(reuse=False))['state'] self.assertFalse(next_answer['cache_review']['hit']) self.assertEqual(self.mocks[6].call_count, 2) async def test_approval_binds_immutable_old_checkpoint(self): first = (await self.ask())['state'] await self.ask('人员有多少?') await api.approve_answer('t1', self.body(first)) record = await self.cache.lookup(first['cache_request']) self.assertEqual(record['payload']['answer_id'], first['answer_id']) async def test_wrong_answer_or_thread_rejected(self): first = (await self.ask())['state'] wrong = self.body(first).model_copy(update={'answer_id': 'forged'}) for thread, body in [('t1', wrong), ('another', self.body(first))]: with self.assertRaises(HTTPException): await api.approve_answer(thread, body) self.assertEqual(self.redis.writes, 0) async def test_stale_approval_rejected(self): first = (await self.ask())['state'] with patch.object(api, 'release_snapshot', return_value={**RELEASE, 'data_version': 'v2'}): with self.assertRaises(HTTPException) as exc: await api.approve_answer('t1', self.body(first)) self.assertEqual(exc.exception.status_code, 409) self.assertEqual(self.redis.writes, 0) async def test_manual_plan_confirmation_is_not_answer_approval(self): event = await self.ask(auto=False) self.assertEqual(event['type'], 'confirm') self.assertEqual(self.redis.writes, 0) done = [e async for e in api._astream_resume('t1', '确认')][-1]['state'] self.assertTrue(done['cache_review']['eligible']) self.assertEqual(self.redis.writes, 0) await api.approve_answer('t1', self.body(done)) self.assertEqual(self.redis.writes, 1) async def test_client_cannot_supply_answer_content(self): from pydantic import ValidationError with self.assertRaises(ValidationError): api.AnswerReviewRequest(answer_id='a', checkpoint_id='b', answer='forged') async def test_write_failure_does_not_report_success(self): first = (await self.ask())['state'] with patch.object(self.redis, 'set', side_effect=TimeoutError): with self.assertRaises(HTTPException) as exc: await api.approve_answer('t1', self.body(first)) self.assertEqual(exc.exception.status_code, 503) async def test_modified_plan_stays_ineligible_after_confirmation(self): await self.ask(auto=False) edits = [e async for e in api._astream_resume('t1', '请修改查询条件')] self.assertEqual(edits[-1]['type'], 'confirm') final = [e async for e in api._astream_resume('t1', '确认')][-1]['state'] self.assertTrue(final['cache_modified']) self.assertFalse(final['cache_review']['eligible']) with self.assertRaises(HTTPException): await api.approve_answer('t1', self.body(final)) async def test_policy_change_invalidates_old_approval(self): first = (await self.ask())['state'] with patch.dict('os.environ', {'QA_ANSWER_CACHE_POLICY_VERSION': 'changed'}): with self.assertRaises(HTTPException): await api.approve_answer('t1', self.body(first)) async def test_revoke_race_not_reported_as_success(self): first = (await self.ask())['state'] result = await api.approve_answer('t1', self.body(first)) with patch.object(self.redis, 'eval', return_value=0): with self.assertRaises(HTTPException) as exc: await api.revoke_answer('t1', self.body(first, entry_id=result['entry_id'])) self.assertEqual(exc.exception.status_code, 409) if __name__ == '__main__': unittest.main(verbosity=2)