231 lines
12 KiB
Python
231 lines
12 KiB
Python
import json
|
|
import math
|
|
|
|
import httpx
|
|
import pytest
|
|
from fastapi.testclient import TestClient
|
|
|
|
from atelier.app import create_app
|
|
|
|
|
|
def fake_provider(request):
|
|
if request.url.path == '/v1/models':
|
|
return httpx.Response(200, json={'data': [{'id': 'local-chat'}, {'id': 'local-embed'}]})
|
|
data = json.loads(request.content)
|
|
if request.url.path == '/v1/embeddings':
|
|
vectors = []
|
|
for i, text in enumerate(data['input']):
|
|
# Known semantic equivalents intentionally have no keyword overlap.
|
|
vectors.append({'index': i, 'embedding': [1., .01] if any(w in text.lower() for w in ['kündigung', 'arbeitsverhältnis', 'job beenden']) else [.01, 1.]})
|
|
return httpx.Response(200, json={'data': vectors})
|
|
if request.url.path == '/v1/chat/completions':
|
|
assert data['model'] == 'local-chat'
|
|
if 'Create archive metadata' in data['messages'][0]['content']:
|
|
assert 'Always write title, description, category and tags in German' in data['messages'][0]['content']
|
|
return httpx.Response(200, json={'choices': [{'message': {'content': json.dumps({'title': 'Kundenanfrage höflich absagen', 'category': 'Kommunikation', 'tags': ['Kunden'], 'description': 'Höfliche Absage'})}}]})
|
|
return httpx.Response(200, json={'choices': [{'message': {'content': 'Verbessert: {{Kunde}}'}}]})
|
|
return httpx.Response(404)
|
|
|
|
|
|
@pytest.fixture
|
|
def client(tmp_path):
|
|
app = create_app(tmp_path, httpx.MockTransport(fake_provider))
|
|
with TestClient(app) as client:
|
|
client.admin_token = (tmp_path / 'admin-token').read_text()
|
|
client.mcp_token = (tmp_path / 'mcp-token').read_text()
|
|
client.post('/api/login', json={'token': client.admin_token})
|
|
yield client
|
|
|
|
|
|
def make(client, title='Arbeitsverhältnis', body='Beende meinen Vertrag höflich.', **kwargs):
|
|
response = client.post('/api/prompts', json=dict(title=title, body=body, **kwargs))
|
|
assert response.status_code == 200, response.text
|
|
return response.json()['prompt']
|
|
|
|
|
|
def configure(client):
|
|
r = client.put('/api/settings', json={'base_url': 'http://model.test/v1', 'api_key': 'secret-model-key', 'chat_model': 'local-chat', 'embedding_model': 'local-embed'})
|
|
assert r.status_code == 200
|
|
assert 'secret-model-key' not in r.text
|
|
|
|
|
|
def test_auth_and_origin(client):
|
|
client.post('/api/logout')
|
|
assert client.get('/api/prompts').status_code == 401
|
|
assert client.post('/api/login', json={'token': 'wrong'}).status_code == 401
|
|
assert client.post('/api/login', json={'token': client.admin_token}, headers={'Origin': 'https://evil.example'}).status_code == 403
|
|
assert client.get('/health').json() == {'status': 'ok'}
|
|
assert client.get('/').headers['content-security-policy']
|
|
|
|
|
|
def test_versions_conflict_trash_restore(client):
|
|
p = make(client)
|
|
updated = dict(p, body='Neue Formulierung', note='Präzisiert')
|
|
r = client.put('/api/prompts/'+p['id'], json=updated)
|
|
assert r.json()['prompt']['version'] == 2
|
|
assert client.put('/api/prompts/'+p['id'], json=updated).status_code == 409
|
|
history = client.get('/api/prompts/'+p['id']+'/versions').json()
|
|
assert len(history) == 2 and history[1]['data']['body'] == p['body']
|
|
client.post('/api/prompts/'+p['id']+'/trash')
|
|
assert client.get('/api/prompts').json()['items'] == []
|
|
assert len(client.get('/api/prompts?trash=true').json()['items']) == 1
|
|
client.post('/api/prompts/'+p['id']+'/trash?deleted=false')
|
|
assert len(client.get('/api/prompts').json()['items']) == 1
|
|
|
|
|
|
def test_semantic_search_and_model_change(client):
|
|
configure(client)
|
|
relevant = make(client)
|
|
make(client, 'Rezept', 'Backe einen Kuchen.')
|
|
assert client.get('/api/prompts?q=Kündigung&mode=text').json()['items'] == []
|
|
result = client.get('/api/prompts?q=Kündigung&mode=semantic').json()
|
|
assert result['items'][0]['id'] == relevant['id'] and result['mode'] == 'semantic'
|
|
assert client.get('/api/settings').json()['indexed'] == 2
|
|
client.put('/api/settings', json={'base_url': 'http://model.test/v1', 'embedding_model': 'different'})
|
|
assert client.get('/api/settings').json()['indexed'] == 0
|
|
assert client.get('/api/prompts?q=Kündigung&mode=semantic').status_code == 502
|
|
assert client.get('/api/prompts?q=Vertrag').json()['warning']
|
|
assert client.post('/api/index').json()['indexed'] == 2
|
|
|
|
|
|
def test_models_improvement_and_secret_preservation(client):
|
|
configure(client)
|
|
assert client.get('/api/models').json()['models'] == ['local-chat', 'local-embed']
|
|
result = client.post('/api/improve', json={'body': 'Hallo {{Kunde}}'}).json()
|
|
assert '{{Kunde}}' in result['body']
|
|
assert client.get('/api/prompts').json()['items'] == []
|
|
client.put('/api/settings', json={'base_url': 'http://model.test/v1', 'api_key': None})
|
|
assert client.get('/api/settings').json()['api_key_set']
|
|
client.put('/api/settings', json={'api_key': ''})
|
|
assert not client.get('/api/settings').json()['api_key_set']
|
|
assert client.put('/api/settings', json={'base_url':'http://user:password@host/v1'}).status_code == 422
|
|
|
|
|
|
def test_export_import_transaction_and_history(client):
|
|
p = make(client)
|
|
client.put('/api/prompts/'+p['id'], json=dict(p, body='v2'))
|
|
archive = client.get('/api/export').json()
|
|
assert 'api_key' not in json.dumps(archive)
|
|
assert client.post('/api/import', json=archive).json()['imported'] == 1
|
|
copies = client.get('/api/prompts').json()['items']
|
|
assert len(copies) == 2 and copies[0]['id'] != copies[1]['id']
|
|
assert len(client.get('/api/prompts/'+copies[0]['id']+'/versions').json()) == 2
|
|
archive['prompts'].append({'title':'broken'})
|
|
assert client.post('/api/import', json=archive).status_code == 422
|
|
assert len(client.get('/api/prompts').json()['items']) == 2
|
|
|
|
|
|
def test_mcp_protocol_and_separate_token(client):
|
|
p = make(client)
|
|
headers={'Authorization':'Bearer '+client.mcp_token, 'Accept':'application/json, text/event-stream'}
|
|
assert client.post('/mcp/', json={}).status_code == 401
|
|
assert client.post('/mcp/', json={}, headers={**headers, 'Authorization':'Bearer '+client.admin_token}).status_code == 401
|
|
r = client.post('/mcp/', headers=headers, json={'jsonrpc':'2.0','id':1,'method':'initialize','params':{'protocolVersion':'2025-03-26','capabilities':{},'clientInfo':{'name':'test','version':'1'}}})
|
|
assert r.status_code == 200, r.text
|
|
assert r.json()['result']['serverInfo']['name'] == 'CasaDePrompt'
|
|
r = client.post('/mcp/', headers=headers, json={'jsonrpc':'2.0','id':2,'method':'tools/list','params':{}})
|
|
assert {t['name'] for t in r.json()['result']['tools']} == {'search_prompts','get_prompt','list_categories','get_prompt_versions'}
|
|
r = client.post('/mcp/', headers=headers, json={'jsonrpc':'2.0','id':3,'method':'tools/call','params':{'name':'search_prompts','arguments':{'query':'Arbeitsverhältnis'}}})
|
|
assert p['id'] in r.text
|
|
client.post('/api/prompts/'+p['id']+'/trash')
|
|
r = client.post('/mcp/', headers=headers, json={'jsonrpc':'2.0','id':4,'method':'tools/call','params':{'name':'get_prompt','arguments':{'prompt_id':p['id']}}})
|
|
assert r.json()['result']['isError']
|
|
|
|
|
|
def test_provider_failure_does_not_lose_prompt(tmp_path):
|
|
transport = httpx.MockTransport(lambda request: httpx.Response(401, text='secret upstream body'))
|
|
with TestClient(create_app(tmp_path, transport)) as c:
|
|
c.post('/api/login', json={'token':(tmp_path/'admin-token').read_text()})
|
|
configure(c)
|
|
response=c.post('/api/prompts', json={'title':'Wichtig','body':'Bleibt erhalten'})
|
|
assert response.status_code == 200 and response.json()['warning']
|
|
assert 'secret upstream body' not in response.text
|
|
assert len(c.get('/api/prompts').json()['items']) == 1
|
|
|
|
|
|
def test_organization_is_a_draft(client):
|
|
configure(client)
|
|
response = client.post('/api/organize', json={'body': 'Absage an {{Kunde}}'})
|
|
assert response.status_code == 200
|
|
assert response.json()['category'] == 'Kommunikation'
|
|
assert client.get('/api/prompts').json()['items'] == []
|
|
|
|
|
|
def test_stale_vectors_do_not_overwrite_new_version(client):
|
|
p = make(client)
|
|
store = client.app.state.store
|
|
updated = store.save({k: p[k] for k in ('title','body','description','category','tags','favorite')}, p['id'], 1)
|
|
store.put_vector(p['id'], 1, 'test', [1, 0])
|
|
assert store.vectors('test') == {}
|
|
store.put_vector(p['id'], updated['version'], 'test', [1, 0])
|
|
assert p['id'] in store.vectors('test')
|
|
|
|
|
|
def test_recheck_endpoint_keeps_draft_and_requires_auth(client):
|
|
configure(client)
|
|
response=client.post('/api/recheck',json={'body':'Original {{name}}','draft':'Draft {{user}}'})
|
|
assert response.status_code==200
|
|
assert response.json()['body']=='Draft {{user}}'
|
|
assert len(response.json()['review']['literal_issues'])==2
|
|
assert response.json()['review']['status']=='unchecked'
|
|
assert client.get('/api/prompts').json()['items']==[]
|
|
client.post('/api/logout')
|
|
assert client.post('/api/recheck',json={'body':'original','draft':'draft'}).status_code==401
|
|
|
|
|
|
def test_german_metadata_is_a_proposal_only(client):
|
|
configure(client)
|
|
before=client.get('/api/prompts').json()
|
|
response=client.post('/api/organize',json={'body':'Write a polite rejection email. Reply in English.'})
|
|
assert response.status_code == 200
|
|
assert response.json()['title'] == 'Kundenanfrage höflich absagen'
|
|
assert set(response.json()) == {'title','description','category','tags'}
|
|
assert client.get('/api/prompts').json() == before
|
|
|
|
|
|
@pytest.mark.parametrize('configured', ['ascii-test-token', 'synthetic-ü-token-🔐'])
|
|
def test_mcp_non_ascii_tokens_never_crash(tmp_path, monkeypatch, configured):
|
|
monkeypatch.setenv('MCP_TOKEN', configured)
|
|
with TestClient(create_app(tmp_path)) as c:
|
|
initialize={'jsonrpc':'2.0','id':1,'method':'initialize','params':{'protocolVersion':'2025-03-26','capabilities':{},'clientInfo':{'name':'regression','version':'1'}}}
|
|
accept={'Accept':'application/json, text/event-stream'}
|
|
for path in ['/mcp', '/mcp/']:
|
|
for wrong in [b'', b'Bearer wrong', 'Bearer falsch-ä'.encode('utf-8'), b'Bearer \xff']:
|
|
assert c.post(path,headers={**accept,'Authorization':wrong},json=initialize).status_code == 401
|
|
headers={**accept,'Authorization':('Bearer '+configured).encode('utf-8')}
|
|
response=c.post(path,headers=headers,json=initialize)
|
|
assert response.status_code == 200
|
|
assert response.json()['result']['serverInfo']['name'] == 'CasaDePrompt'
|
|
listed=c.post(path,headers=headers,json={'jsonrpc':'2.0','id':2,'method':'tools/list','params':{}})
|
|
assert len(listed.json()['result']['tools']) == 4
|
|
|
|
|
|
def test_web_login_with_unicode_token(tmp_path, monkeypatch):
|
|
monkeypatch.setenv('ADMIN_TOKEN', 'synthetic-ä-admin-🔐')
|
|
with TestClient(create_app(tmp_path)) as c:
|
|
assert c.post('/api/login',json={'token':'falsch-ü'}).status_code == 401
|
|
assert c.post('/api/login',json={'token':'synthetic-ä-admin-🔐'}).status_code == 200
|
|
|
|
|
|
def test_mcp_id_aliases_and_conflicts(client):
|
|
p=make(client)
|
|
headers={'Authorization':'Bearer '+client.mcp_token,'Accept':'application/json, text/event-stream'}
|
|
def call(tool,args):
|
|
response=client.post('/mcp/',headers=headers,json={'jsonrpc':'2.0','id':1,'method':'tools/call','params':{'name':tool,'arguments':args}})
|
|
assert response.status_code == 200
|
|
return response.json()['result']
|
|
listed=client.post('/mcp/',headers=headers,json={'jsonrpc':'2.0','id':2,'method':'tools/list','params':{}}).json()['result']['tools']
|
|
for tool in ['get_prompt','get_prompt_versions']:
|
|
schema=next(t['inputSchema'] for t in listed if t['name']==tool)
|
|
assert {'id','prompt_id'} <= set(schema['properties'])
|
|
canonical=call(tool,{'prompt_id':p['id']})
|
|
assert not canonical.get('isError')
|
|
assert call(tool,{'id':p['id']}) == canonical
|
|
assert call(tool,{'id':p['id'],'prompt_id':p['id']}) == canonical
|
|
for args in [{},{'id':''},{'id':' '},{'id':p['id'],'prompt_id':'different'}]:
|
|
assert call(tool,args)['isError']
|
|
assert call('get_prompt',{'id':p['id'],'version':1}) == call('get_prompt',{'prompt_id':p['id'],'version':1})
|
|
client.post('/api/prompts/'+p['id']+'/trash')
|
|
assert call('get_prompt',{'id':p['id']})['isError']
|
|
assert call('get_prompt_versions',{'id':p['id']})['isError']
|