190 lines
7 KiB
Python
190 lines
7 KiB
Python
"""M8: user edits to transcript/summary via PUT endpoints.
|
|
|
|
Edits create a new version with edited_by_user=True; the AI pipeline
|
|
must not overwrite them afterwards.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import struct
|
|
import uuid
|
|
|
|
from sqlalchemy import select
|
|
|
|
from shonar.db import session as db_session
|
|
from shonar.db.models import Transcript
|
|
from shonar.services import processing
|
|
|
|
AUTH = {"email": "m8@example.com", "password": "m8-test-passw0rd-123"}
|
|
|
|
|
|
async def user_tokens(client):
|
|
r = await client.post(
|
|
"/api/v1/auth/register",
|
|
json={"email": AUTH["email"], "password": AUTH["password"]},
|
|
)
|
|
assert r.status_code == 201, r.text
|
|
return r.json()["access_token"]
|
|
|
|
|
|
def wav_bytes(payload_len: int = 64) -> bytes:
|
|
data = bytes(range(payload_len % 256)) * (payload_len // 256 + 1)
|
|
data = data[:payload_len]
|
|
header = (
|
|
b"RIFF" + struct.pack("<I", 36 + len(data)) + b"WAVE"
|
|
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, 8000, 8000, 1, 8)
|
|
+ b"data" + struct.pack("<I", len(data))
|
|
)
|
|
return header + data
|
|
|
|
|
|
async def upload_recording(client, token, client_id=None):
|
|
h = {"Authorization": f"Bearer {token}"}
|
|
data = wav_bytes()
|
|
r = await client.post(
|
|
"/api/v1/uploads",
|
|
json={"declared_mime_type": "audio/wav", "declared_size_bytes": len(data),
|
|
"client_recording_id": client_id},
|
|
headers=h,
|
|
)
|
|
assert r.status_code == 201, r.text
|
|
sid = r.json()["id"]
|
|
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
|
headers={**h, "content-type": "application/octet-stream"})
|
|
assert r.status_code == 201, r.text
|
|
r = await client.post(f"/api/v1/uploads/{sid}/finalize",
|
|
json={"duration_seconds": 5.0}, headers=h)
|
|
assert r.status_code == 201, r.text
|
|
return r.json()
|
|
|
|
|
|
async def test_put_transcript_creates_user_version(client):
|
|
token = await user_tokens(client)
|
|
rec = await upload_recording(client, token, client_id="m8-t-1")
|
|
h = {"Authorization": f"Bearer {token}"}
|
|
rid = rec["id"]
|
|
|
|
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
|
json={"text": "hello corrected"}, headers=h)
|
|
assert r.status_code == 200, r.text
|
|
body = r.json()
|
|
assert body["version"] == 1
|
|
assert body["text"] == "hello corrected"
|
|
assert body["edited_by_user"] is True
|
|
assert body["provider"] == "user"
|
|
|
|
r = await client.get(f"/api/v1/recordings/{rid}/transcript", headers=h)
|
|
assert r.status_code == 200
|
|
assert r.json()["text"] == "hello corrected"
|
|
|
|
# Second edit supersedes the first.
|
|
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
|
json={"text": "second pass",
|
|
"segments": [{"start": 0.0, "end": 1.0, "text": "second pass"}]},
|
|
headers=h)
|
|
assert r.status_code == 200, r.text
|
|
assert r.json()["version"] == 2
|
|
assert r.json()["segments"][0]["text"] == "second pass"
|
|
|
|
async with db_session._session_factory() as s:
|
|
rows = (await s.scalars(
|
|
select(Transcript).where(Transcript.recording_id == uuid.UUID(rid))
|
|
.order_by(Transcript.version))).all()
|
|
assert [(r.version, r.superseded_at is not None, r.edited_by_user) for r in rows] == [
|
|
(1, True, True), (2, False, True)]
|
|
|
|
|
|
async def test_put_summary_creates_user_version(client):
|
|
token = await user_tokens(client)
|
|
rec = await upload_recording(client, token, client_id="m8-s-1")
|
|
h = {"Authorization": f"Bearer {token}"}
|
|
rid = rec["id"]
|
|
|
|
content = {"short": "s", "action_items": ["ship it"]}
|
|
r = await client.put(f"/api/v1/recordings/{rid}/summary",
|
|
json={"content": content}, headers=h)
|
|
assert r.status_code == 200, r.text
|
|
body = r.json()
|
|
assert body["version"] == 1
|
|
assert body["content"] == content
|
|
assert body["edited_by_user"] is True
|
|
|
|
r = await client.get(f"/api/v1/recordings/{rid}/summary", headers=h)
|
|
assert r.status_code == 200
|
|
assert r.json()["content"]["action_items"] == ["ship it"]
|
|
|
|
|
|
async def test_user_edit_survives_auto_pipeline(client, monkeypatch):
|
|
"""An auto transcribe run after a user edit inserts nothing."""
|
|
from shonar.services.ai import Segment as AiSegment
|
|
from shonar.services.ai import SummaryResult, TranscriptResult
|
|
|
|
class FakeTranscriber:
|
|
name = "fake-stt"
|
|
|
|
async def transcribe(self, audio, mime, *, language_hint=None):
|
|
return TranscriptResult(
|
|
text="auto text", language="en",
|
|
segments=[AiSegment(0.0, 1.0, "auto text")], model="fake-stt-1",
|
|
)
|
|
|
|
class FakeLlm:
|
|
name = "fake-llm"
|
|
|
|
async def summarize(self, transcript, *, title=None, tone=None,
|
|
on_progress=None):
|
|
return SummaryResult(short="s", detailed="d", key_points=(),
|
|
decisions=(), action_items=(), questions=(),
|
|
model="fake-llm-1")
|
|
|
|
monkeypatch.setattr(
|
|
processing, "get_transcription_provider", lambda settings: FakeTranscriber())
|
|
monkeypatch.setattr(processing, "get_llm_provider", lambda settings: FakeLlm())
|
|
|
|
token = await user_tokens(client)
|
|
rec = await upload_recording(client, token, client_id="m8-t-2")
|
|
h = {"Authorization": f"Bearer {token}"}
|
|
rid = rec["id"]
|
|
|
|
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
|
json={"text": "user verdict"}, headers=h)
|
|
assert r.status_code == 200, r.text
|
|
|
|
monkeypatch.setattr(
|
|
processing, "get_transcription_provider", lambda settings: FakeTranscriber())
|
|
monkeypatch.setattr(processing, "get_llm_provider", lambda settings: FakeLlm())
|
|
await processing.run_transcribe({}, rid)
|
|
|
|
r = await client.get(f"/api/v1/recordings/{rid}/transcript", headers=h)
|
|
assert r.status_code == 200
|
|
assert r.json()["text"] == "user verdict"
|
|
|
|
|
|
async def test_edit_endpoints_enforce_ownership(client):
|
|
token = await user_tokens(client)
|
|
rec = await upload_recording(client, token, client_id="m8-t-3")
|
|
rid = rec["id"]
|
|
|
|
r = await client.post("/api/v1/auth/register",
|
|
json={"email": "m8-other@example.com", "password": "m8-other-passw0rd-1"})
|
|
assert r.status_code == 201, r.text
|
|
other = r.json()["access_token"]
|
|
h2 = {"Authorization": f"Bearer {other}"}
|
|
|
|
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
|
json={"text": "hijack"}, headers=h2)
|
|
assert r.status_code == 404
|
|
r = await client.put(f"/api/v1/recordings/{rid}/summary",
|
|
json={"content": {}}, headers=h2)
|
|
assert r.status_code == 404
|
|
|
|
|
|
async def test_edit_validation(client):
|
|
token = await user_tokens(client)
|
|
rec = await upload_recording(client, token, client_id="m8-t-4")
|
|
h = {"Authorization": f"Bearer {token}"}
|
|
rid = rec["id"]
|
|
|
|
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
|
json={"text": ""}, headers=h)
|
|
assert r.status_code == 422
|