| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299 |
- """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_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)
|