feat: improve chat retrieval, message versions and Markdown rendering

This commit is contained in:
2026-09-06 21:39:40 +08:00
parent 174e723545
commit 637ddbb9bf
34 changed files with 1072 additions and 62 deletions
+4 -9
View File
@@ -11,7 +11,7 @@ from app.services.chat_context import prepare
@pytest.mark.parametrize('enabled', [True, False])
def test_chat_stream_retrieves_real_notes_and_emits_sources(monkeypatch, enabled):
def test_chat_stream_does_not_presearch_notes(monkeypatch, enabled):
received = []
class Adapter:
@@ -34,14 +34,9 @@ def test_chat_stream_retrieves_real_notes_and_emits_sources(monkeypatch, enabled
assert [e['sequence'] for e in events] == list(range(len(events)))
assert events[-1]['event'] == 'Done'
assert received[0].messages == request.messages
if enabled:
assert events[0]['event'] == 'Citation'
assert events[0]['data']['note_id'] == note.note_id
assert 'apple orchard knowledge' in received[0].system
assert 'Keep original instructions' in received[0].system
else:
assert all(e['event'] != 'Citation' for e in events)
assert received[0].system == request.system
assert all(e['event'] != 'Citation' for e in events)
assert 'apple orchard knowledge' not in received[0].system
assert 'Keep original instructions' in received[0].system
assert request.system == 'Keep original instructions'
asyncio.run(scenario())
+143
View File
@@ -0,0 +1,143 @@
import asyncio
from types import SimpleNamespace
import pytest
from app.contracts import ChatRequest, Message, ModelCapability, ModelEventType as E
from app.services import chat_retrieval as service
def test_stream_searches_again_and_preserves_numbers(monkeypatch):
seen = []
async def prepare(request):
query = request.retrieval.query if request.retrieval else 'initial'
return request, [{'block_id': 'a' if query == 'initial' else 'b', 'number': 1, 'content': query, 'citation_id': 'cit_blk_test'}]
monkeypatch.setattr(service, 'prepare', prepare)
class Adapter:
async def stream(self, request):
seen.append(request)
if len(seen) == 1:
yield service.event(E.text_delta, {'text': '需要补充资料。'})
yield service.event(E.tool_call_start, {'tool_call_id': 'call', 'name': 'rag.search'})
yield service.event(E.tool_call_delta, {'tool_call_id': 'call', 'arguments_delta': '{"query":"new"}'})
yield service.event(E.tool_call_end, {'tool_call_id': 'call'})
else:
assert request.messages[-1].role.value == 'tool'
assert '"number": 1' in request.messages[-1].content
assert 'cit_blk_test' not in request.messages[-1].content
assert 'block_id' not in request.messages[-1].content
yield service.event(E.text_delta, {'text': '根据新证据 [1]'})
yield service.event(E.usage, {'input_tokens': 10, 'output_tokens': 2})
yield service.event(E.done, {})
provider = SimpleNamespace(adapter=Adapter(), config=SimpleNamespace(capabilities=[ModelCapability.tool_calling]))
request = ChatRequest(provider_id='x', model='x', messages=[Message(role='user', content='question')])
async def run(): return [item async for item in service.stream(request, provider)]
events = asyncio.run(run())
assert len(seen) == 2
assert any(e.event == E.text_delta and e.data['text'] == '\n\n' for e in events)
assert events[0].event == E.text_delta
assert [e.data['number'] for e in events if e.event == E.citation] == [1]
assert sum(e.event == E.done for e in events) == 1
assert next(e.data for e in events if e.event == E.usage) == {'input_tokens': 20, 'output_tokens': 4}
assert [e.event for e in events].index(E.tool_call_end) > max(i for i, e in enumerate(events) if e.event == E.citation)
@pytest.mark.parametrize('tool_name', ['rag.search', 'notes.update'])
def test_loop_is_bounded_and_never_executes_write_tools(monkeypatch, tool_name):
searches, requests = [], []
async def prepare(request):
searches.append(request)
return request, []
monkeypatch.setattr(service, 'prepare', prepare)
class Adapter:
async def stream(self, request):
requests.append(request)
yield service.event(E.tool_call_start, {'tool_call_id': 'same', 'name': tool_name, 'arguments': {'query': 'again'}})
yield service.event(E.done, {})
provider = SimpleNamespace(adapter=Adapter(), config=SimpleNamespace(capabilities=[ModelCapability.tool_calling]))
async def run():
return [e async for e in service.stream(ChatRequest(provider_id='x', model='x', messages=[Message(role='user', content='q')]), provider)]
events = asyncio.run(run())
assert len(requests) == 4
assert requests[-1].tools == []
assert len(searches) == (3 if tool_name == 'rag.search' else 0)
assert len({e.data['tool_call_id'] for e in events if e.event == E.tool_call_start}) == 4
assert events[-1].data['status'] == 'failed'
def test_closing_stream_closes_provider(monkeypatch):
closed = []
async def prepare(request): return request, []
monkeypatch.setattr(service, 'prepare', prepare)
class Adapter:
async def stream(self, request):
try:
yield service.event(E.text_delta, {'text': 'partial'})
await asyncio.sleep(60)
finally:
closed.append(True)
async def run():
provider = SimpleNamespace(adapter=Adapter(), config=SimpleNamespace(capabilities=[ModelCapability.tool_calling]))
events = service.stream(ChatRequest(provider_id='x', model='x', messages=[Message(role='user', content='q')]), provider)
await anext(events)
await events.aclose()
asyncio.run(run())
assert closed == [True]
def test_no_search_without_a_model_call_and_timeout_allows_continuation(monkeypatch):
called = []
monkeypatch.setattr(service, 'SEARCH_TIMEOUT_SECONDS', .01)
async def slow_search(request):
called.append(True)
await asyncio.sleep(10)
monkeypatch.setattr(service, 'prepare', slow_search)
requests = []
class Adapter:
async def stream(self, request):
requests.append(request)
if len(requests) == 1:
assert called == []
yield service.event(E.text_delta, {'text': '我来查看笔记。'})
yield service.event(E.tool_call_start, {'tool_call_id': 'search', 'name': 'rag.search', 'arguments': {'query': 'q'}})
else:
assert 'Retrieval failed' in request.messages[-1].content
yield service.event(E.text_delta, {'text': '检索超时,暂时无法核对笔记。'})
yield service.event(E.done, {})
async def run():
provider = SimpleNamespace(adapter=Adapter(), config=SimpleNamespace(capabilities=[ModelCapability.tool_calling]))
return [e async for e in service.stream(ChatRequest(provider_id='x', model='x', messages=[Message(role='user', content='q')]), provider)]
events = asyncio.run(run())
assert events[0].event == E.text_delta
assert next(e for e in events if e.event == E.tool_call_end).data['status'] == 'failed'
assert events[-1].data['status'] == 'completed'
def test_thinking_is_replayed_on_real_compatible_wire(monkeypatch):
import json
import httpx
from app.providers.openai_compatible import OpenAICompatibleProvider
requests = []
async def prepare(request): return request, []
monkeypatch.setattr(service, 'prepare', prepare)
def handler(request):
payload = json.loads(request.content)
requests.append(payload)
if len(requests) == 1:
alias = payload['tools'][0]['function']['name']
deltas = [{'reasoning_content': 'Need '}, {'reasoning_content': 'more evidence.'},
{'tool_calls': [{'index': i, 'id': f'call{i}', 'type': 'function', 'function': {'name': alias, 'arguments': '{"query":"Python"}'}} for i in range(2)]}]
else:
assistant = next(m for m in payload['messages'] if m.get('tool_calls'))
if assistant.get('reasoning_content') != 'Need more evidence.':
return httpx.Response(400, json={'error': {'message': 'reasoning_content required'}})
assert {c['id'] for c in assistant['tool_calls']} == {m['tool_call_id'] for m in payload['messages'] if m['role'] == 'tool'}
deltas = [{'content': 'Answer after retrieval'}]
body = ''.join('data: ' + json.dumps({'choices': [{'delta': delta}]}) + '\n\n' for delta in deltas) + 'data: [DONE]\n\n'
return httpx.Response(200, text=body, headers={'content-type': 'text/event-stream'})
adapter = OpenAICompatibleProvider('https://provider.test', None, SimpleNamespace(resolve=lambda _: None), transport=httpx.MockTransport(handler))
provider = SimpleNamespace(adapter=adapter, config=SimpleNamespace(capabilities=[ModelCapability.tool_calling]))
async def run():
return [e async for e in service.stream(ChatRequest(provider_id='x', model='x', messages=[Message(role='user', content='q')]), provider)]
events = asyncio.run(run())
assert len(requests) == 2
assert not any(e.event == E.error for e in events)
assert any(e.data.get('text') == 'Answer after retrieval' for e in events)
+40
View File
@@ -0,0 +1,40 @@
from app.services import chat_history as history
def test_edits_regeneration_and_activity_survive_version_switch():
history.create('Versions', 'versions')
def append(id, role, content, parent=None, activity=None):
history.append_message('versions', message_id=id, role=role, content=content, parent_message_id=parent, activity=activity)
append('u1', 'user', 'original')
append('a1', 'assistant', 'original answer', 'u1')
append('u2', 'user', 'follow-up')
append('a2', 'assistant', 'follow-up answer', 'u2')
history.prepare_retry('versions', 'u1')
append('u1-edit', 'user', 'edited')
history.reserve_response('versions', 'a1-edit')
trace = [{'type': 'thinking', 'text': 'before'}, {'type': 'tool', 'tool_call_id': 'tool'}, {'type': 'thinking', 'text': 'after'}]
append('a1-edit', 'assistant', 'edited answer', 'u1-edit', trace)
items, _ = history.list_messages('versions', 500, 0)
assert [m.message_id for m in items] == ['u1-edit', 'a1-edit']
assert items[0].versions == ['u1', 'u1-edit']
assert items[1].activity == trace
history.select_version('versions', 'u1')
assert [m.message_id for m in history.list_messages('versions', 500, 0)[0]] == ['u1', 'a1', 'u2', 'a2']
history.prepare_retry('versions', 'a1')
history.reserve_response('versions', 'a1-new')
append('a1-new', 'assistant', 'regenerated', 'u1')
items, _ = history.list_messages('versions', 500, 0)
assert [m.message_id for m in items] == ['u1', 'a1-new']
assert items[-1].versions == ['a1', 'a1-new']
history.select_version('versions', 'a1')
assert history.list_messages('versions', 500, 0)[0][-1].message_id == 'a2'
def test_late_response_does_not_replace_new_generation():
history.create('Late', 'late')
history.append_message('late', message_id='u', role='user', content='question')
history.reserve_response('late', 'new')
history.append_message('late', message_id='old', role='assistant', content='old', parent_message_id='u')
assert history.list_messages('late', 500, 0)[0][-1].message_id == 'u'
history.append_message('late', message_id='new', role='assistant', content='new', parent_message_id='u')
assert history.list_messages('late', 500, 0)[0][-1].message_id == 'new'
+45
View File
@@ -0,0 +1,45 @@
import asyncio
import hashlib
from typing import get_args
import pytest
from app.agent.markdown_tools import ComposeArguments, Format, PatchArguments, compose, patch, register
from app.agent.tools import ToolRegistry
from app.services import note_service
@pytest.mark.parametrize('kind', get_args(Format))
def test_all_registered_formats_compose(kind):
result = compose(ComposeArguments(format=kind, text='Example', items=['one', 'two'], rows=[['A', 'B'], ['C', 'D']], url='https://example.com', title='Title', tags=['tag']), None)
assert result['markdown']
assert result['persisted'] is False
def test_fences_tables_and_permissions():
assert compose(ComposeArguments(format='code-block', text='```'), None)['markdown'].startswith('````\n')
with pytest.raises(ValueError): compose(ComposeArguments(format='table', rows=[['a'], ['b', 'c']]), None)
registry = ToolRegistry()
register(registry)
assert registry.get('notes.patch_markdown').definition.permission == 'notes.write'
assert registry.get('markdown.compose').definition.permission is None
def test_patch_preserves_unrelated_content_and_rejects_stale_version():
async def run():
note = await note_service.create_note(title='Patch test', markdown='before\n\nold\n\nafter', folder=None, tags=[])
args = PatchArguments(note_id=note.note_id, expected_content_hash=hashlib.sha256(note.markdown.encode()).hexdigest(), old_text='old', new_text='> [!NOTE]\n> new')
await patch(args, None)
updated = await note_service.get_note(note.note_id)
assert updated.markdown == 'before\n\n> [!NOTE]\n> new\n\nafter'
with pytest.raises(ValueError): await patch(args, None)
asyncio.run(run())
def test_metadata_patch_updates_index_tags():
async def run():
markdown = '---\ntitle: Old\ntags: [old]\n---\nBody'
note = await note_service.create_note(title='Old', markdown=markdown, folder=None, tags=[])
await patch(PatchArguments(note_id=note.note_id, expected_content_hash=hashlib.sha256(markdown.encode()).hexdigest(), old_text='tags: [old]', new_text='tags: [new]'), None)
updated = await note_service.get_note(note.note_id)
assert updated.tags == ['new']
assert updated.markdown.endswith('Body')
asyncio.run(run())