161 lines
8.2 KiB
Python
161 lines
8.2 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 'Ordne' in data['messages'][0]['content']:
|
|
return httpx.Response(200, json={'choices': [{'message': {'content': json.dumps({'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')
|