Files
treg/tests/test_fishaudio.py

467 lines
17 KiB
Python

"""Fish Audio shared-key tenancy, lifecycle and pricing contracts."""
from __future__ import annotations
import json
import pytest
from httpx import AsyncClient
from sqlalchemy import select
from treg import audit
from treg.application.call import resolve as call_resolution
from treg.application.call import service as call_service
from treg.application.call.types import UpstreamResponse
from treg.config import get_settings
from treg.domain import provider_resources
from treg.domain.identity import api_keys as managed_keys
from treg.infra.db import _engine, session_maker
from treg.models import ProviderResource
@pytest.fixture
def fishaudio_platform_on(monkeypatch):
monkeypatch.setenv("TREG_PLATFORM_KEY_FISHAUDIO", "PLATFORM-FISH-KEY")
monkeypatch.setenv("TREG_PLATFORM_PROVIDERS", "fishaudio")
get_settings.cache_clear()
yield
get_settings.cache_clear()
def _response(status: int, body: bytes = b"{}", headers: tuple = ()):
async def stream():
yield body
async def close():
return None
return UpstreamResponse(status, headers, stream(), close)
async def _org_id(clients: AsyncClient) -> int:
return (await clients.get("/orgs")).json()[0]["org_id"]
async def _add_voice(clients: AsyncClient, voice_id: str, name: str = "Narrator") -> None:
async with session_maker() as db:
db.add(ProviderResource(
org_id=await _org_id(clients), provider="fishaudio", resource_kind="voice",
upstream_id=voice_id, display_name=name, created_by="tim@superdesign.dev",
source_call_id="call-test",
))
await db.commit()
async def test_voice_create_is_private_and_becomes_an_org_resource(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
seen = []
async def relay(request, upstream_url, tool, secrets, client, **kwargs):
seen.append((request.method, upstream_url))
return _response(201, b'{"_id":"voice-created","title":"Narrator"}')
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.post(
"/call/fishaudio.voices.create",
data={"type": "tts", "title": "Narrator", "train_mode": "fast", "visibility": "private"},
files=[("voices", ("voice.wav", b"RIFF-test", "audio/wav"))],
)
assert response.status_code == 201, response.text
listed = await clients.get(
f"/orgs/{await _org_id(clients)}/provider-resources?provider=fishaudio&kind=voice"
)
assert listed.headers["X-Treg-Resource-Source"] == "platform"
rows = listed.json()
assert [(row["upstream_id"], row["display_name"], row["status"]) for row in rows] == [
("voice-created", "Narrator", "active")
]
assert seen == [("POST", "https://api.fish.audio/model")]
async def test_unified_voice_list_uses_byok_and_normalizes_without_holding_db(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
await clients.post("/secrets", json={"name": "fishaudio", "value": "OWN-FISH-KEY"})
await audit.drain()
checked_out = []
async def relay(*args, **kwargs):
checked_out.append(_engine.pool.checkedout())
return _response(200, b'{"items":[{"_id":"account-voice","title":"Account narrator",'
b'"state":"created","created_at":"2026-09-22T00:00:00Z"}]}')
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.get(
f"/orgs/{await _org_id(clients)}/provider-resources?provider=fishaudio&kind=voice"
)
assert response.status_code == 200, response.text
assert response.headers["X-Treg-Resource-Source"] == "byok"
assert response.json() == [{
"id": "account-voice",
"provider": "fishaudio",
"kind": "voice",
"upstream_id": "account-voice",
"display_name": "Account narrator",
"created_by": None,
"source_call_id": None,
"status": "active",
"created_at": "2026-09-22T00:00:00Z",
"updated_at": None,
"deleted_at": None,
}]
assert checked_out == [0]
async def test_unified_voice_list_follows_fish_pagination(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
await clients.post("/secrets", json={"name": "fishaudio", "value": "OWN-FISH-KEY"})
seen_pages = []
async def relay(request, *args, **kwargs):
page = dict(request.query_items)["page_number"]
seen_pages.append(page)
body = (
b'{"items":[{"_id":"voice-one","title":"One"}],"has_more":true}'
if page == "1" else
b'{"items":[{"_id":"voice-two","title":"Two"}],"has_more":false}'
)
return _response(200, body)
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.get(
f"/orgs/{await _org_id(clients)}/provider-resources?provider=fishaudio&kind=voice"
)
assert response.status_code == 200, response.text
assert [row["upstream_id"] for row in response.json()] == ["voice-one", "voice-two"]
assert seen_pages == ["1", "2"]
async def test_platform_resource_source_never_substitutes_byok_account_rows(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
await _add_voice(clients, "team-voice", "Team narrator")
await clients.post("/secrets", json={"name": "fishaudio", "value": "OWN-FISH-KEY"})
called = False
async def relay(*args, **kwargs):
nonlocal called
called = True
return _response(200, b'{"items":[{"_id":"account-voice"}]}')
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.get(
f"/orgs/{await _org_id(clients)}/provider-resources?source=platform"
)
assert response.status_code == 200, response.text
assert response.headers["X-Treg-Resource-Source"] == "platform"
assert [row["upstream_id"] for row in response.json()] == ["team-voice"]
assert called is False
async def test_voice_create_holds_no_db_connection_during_fish_io(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
await audit.drain()
checked_out = []
async def relay(*args, **kwargs):
checked_out.append(_engine.pool.checkedout())
return _response(201, b'{"_id":"voice-no-held-session"}')
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.post(
"/call/fishaudio.voices.create",
data={"type": "tts", "title": "No held session", "train_mode": "fast",
"visibility": "private"},
files=[("voices", ("voice.wav", b"RIFF-test", "audio/wav"))],
)
assert response.status_code == 201, response.text
assert checked_out == [0]
async def test_platform_voice_create_rejects_non_private_before_relay(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
called = False
async def relay(*args, **kwargs):
nonlocal called
called = True
return _response(201, b'{"_id":"must-not-exist"}')
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.post(
"/call/fishaudio.voices.create",
data={"type": "tts", "title": "Public", "train_mode": "fast", "visibility": "public"},
files=[("voices", ("voice.wav", b"RIFF-test", "audio/wav"))],
)
assert response.status_code == 400
assert called is False
async def test_unknown_and_cross_org_voice_ids_are_identical_403_without_upstream(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
await _add_voice(clients, "voice-team-one")
called = 0
async def relay(*args, **kwargs):
nonlocal called
called += 1
return _response(200)
monkeypatch.setattr(call_service, "relay", relay)
unknown = await clients.patch(
"/call/fishaudio.voices.update?id=unknown",
json={"title": "Nope", "visibility": "private"})
other = (await clients.post("/orgs", json={"name": "Other Fish Team"})).json()
clients.headers["X-Treg-Token"] = other["token"]
clients.headers.pop("X-Treg-Org", None)
cross_org = await clients.patch(
"/call/fishaudio.voices.update?id=voice-team-one",
json={"title": "Nope", "visibility": "private"})
assert unknown.status_code == cross_org.status_code == 403
assert unknown.json()["detail"] == cross_org.json()["detail"]
assert called == 0
async def test_cross_org_tts_voice_is_denied_without_public_lookup(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
await _add_voice(clients, "private-team-one")
other = (await clients.post("/orgs", json={"name": "Other Fish TTS Team"})).json()
clients.headers["X-Treg-Token"] = other["token"]
clients.headers.pop("X-Treg-Org", None)
called = False
async def relay(*args, **kwargs):
nonlocal called
called = True
return _response(200)
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.post(
"/call/fishaudio.tts.s2-1-pro",
headers={"model": "s2.1-pro"},
json={"text": "Must stay private", "reference_id": "private-team-one"},
)
assert response.status_code == 403
assert called is False
async def test_tts_checks_every_voice_in_an_array(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
await _add_voice(clients, "voice-a")
calls = []
async def relay(request, upstream_url, *args, **kwargs):
calls.append((request.method, upstream_url))
return _response(200, b'{"_id":"not-ours","visibility":"private","licensed":false}')
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.post(
"/call/fishaudio.tts.s2-1-pro",
headers={"model": "s2.1-pro"},
json={"text": "Two voices", "reference_id": ["voice-a", "not-ours"]},
)
assert response.status_code == 403
assert calls == [("GET", "https://api.fish.audio/model/not-ours")]
async def test_tts_accepts_a_live_verified_public_community_voice_without_holding_db(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
await audit.drain()
# SQLite tests share the API engine with the best-effort managed-key metadata writer.
# Finish the writer scheduled by fixture login before inspecting the global pool count.
await managed_keys.drain_last_used()
# Auth on this call can schedule another optional display-metadata write. Keep the pool
# assertion scoped to the call runtime whose transaction boundary this test guards.
monkeypatch.setattr(managed_keys, "touch", lambda *args, **kwargs: None)
calls = []
checked_out = []
async def relay(request, upstream_url, *args, **kwargs):
calls.append((request.method, upstream_url))
checked_out.append(_engine.pool.checkedout())
if request.method == "GET":
return _response(
200,
b'{"_id":"public-voice","visibility":"public","licensed":false}',
)
return _response(200, b"audio", ((b"content-type", b"audio/mpeg"),))
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.post(
"/call/fishaudio.tts.s2-1-pro",
headers={"model": "s2.1-pro"},
json={"text": "Public voice", "reference_id": "public-voice"},
)
assert response.status_code == 200, response.text
assert response.content == b"audio"
assert calls == [
("GET", "https://api.fish.audio/model/public-voice"),
("POST", "https://api.fish.audio/v1/tts"),
]
assert checked_out == [0, 0]
@pytest.mark.parametrize("licensed", ["true", "false"])
async def test_public_voice_discovery_is_a_free_platform_action(
clients: AsyncClient, fishaudio_platform_on, monkeypatch, licensed,
):
await audit.drain()
observed = []
async def relay(request, upstream_url, *args, **kwargs):
observed.append((upstream_url, dict(request.query_items), _engine.pool.checkedout()))
return _response(
200,
b'{"total":1,"items":[{"_id":"public-voice","visibility":"public",'
b'"licensed":true}],"has_more":false}',
((b"content-type", b"application/json"),),
)
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.get(
"/call/fishaudio.voices.discover",
params={"self": "false", "licensed": licensed, "page_size": "3", "page_number": "1"},
)
assert response.status_code == 200, response.text
assert response.json()["items"][0]["_id"] == "public-voice"
assert observed == [(
"https://api.fish.audio/model",
{"self": "false", "licensed": licensed, "page_size": "3", "page_number": "1"},
0,
)]
@pytest.mark.parametrize("reference_id", [None, "voice-a", ["voice-a", "voice-b"]])
async def test_tts_accepts_default_scalar_and_owned_array_references(
clients: AsyncClient, fishaudio_platform_on, monkeypatch, reference_id,
):
await _add_voice(clients, "voice-a")
await _add_voice(clients, "voice-b")
called = 0
async def relay(*args, **kwargs):
nonlocal called
called += 1
return _response(200, b"audio", ((b"content-type", b"audio/mpeg"),))
monkeypatch.setattr(call_service, "relay", relay)
body = {"text": "Owned voice"}
if reference_id is not None:
body["reference_id"] = reference_id
response = await clients.post(
"/call/fishaudio.tts.s2-1-pro", headers={"model": "s2.1-pro"}, json=body,
)
assert response.status_code == 200, response.text
assert response.content == b"audio"
assert called == 1
async def test_platform_model_header_must_be_one_unambiguous_value(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
called = False
async def relay(*args, **kwargs):
nonlocal called
called = True
return _response(200, b"audio")
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.post(
"/call/fishaudio.tts.s2-1-pro",
headers=[("model", "s1"), ("model", "s2.1-pro")],
json={"text": "ambiguous"},
)
assert response.status_code == 400
assert called is False
async def test_byok_bypasses_platform_voice_ownership(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
await clients.post("/secrets", json={"name": "fishaudio", "value": "OWN-FISH-KEY"})
called = 0
async def relay(request, upstream_url, tool, secrets, client, **kwargs):
nonlocal called
called += 1
return _response(200, b'{"_id":"foreign-account-voice"}')
monkeypatch.setattr(call_service, "relay", relay)
response = await clients.patch(
"/call/fishaudio.voices.update?id=foreign-account-voice", json={"title": "Mine"})
assert response.status_code == 200, response.text
assert called == 1
async def test_voice_update_and_delete_change_local_lifecycle_only_after_fish(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
await _add_voice(clients, "voice-life", "Old name")
statuses = iter((200, 404, 404))
async def relay(*args, **kwargs):
return _response(next(statuses))
monkeypatch.setattr(call_service, "relay", relay)
updated = await clients.patch(
"/call/fishaudio.voices.update?id=voice-life",
json={"title": "New name", "visibility": "private"},
)
deleted = await clients.delete("/call/fishaudio.voices.delete?id=voice-life")
retried = await clients.delete("/call/fishaudio.voices.delete?id=voice-life")
assert updated.status_code == 200
assert deleted.status_code == retried.status_code == 404
async with session_maker() as db:
row = (await db.execute(select(ProviderResource).where(
ProviderResource.upstream_id == "voice-life"
))).scalars().one()
assert row.display_name == "New name"
assert row.status == provider_resources.DELETED and row.deleted_at is not None
async def test_create_persistence_failure_deletes_the_unmanaged_fish_voice(
clients: AsyncClient, fishaudio_platform_on, monkeypatch,
):
calls = []
async def relay(request, upstream_url, tool, secrets, client, **kwargs):
calls.append((request.method, upstream_url))
if request.method == "DELETE":
return _response(204, b"")
return _response(201, b'{"_id":"orphan-candidate"}')
async def fail_register(*args, **kwargs):
raise RuntimeError("database unavailable")
monkeypatch.setattr(call_service, "relay", relay)
monkeypatch.setattr(provider_resources, "register", fail_register)
response = await clients.post(
"/call/fishaudio.voices.create",
data={"type": "tts", "title": "Unsafe", "train_mode": "fast", "visibility": "private"},
files=[("voices", ("voice.wav", b"RIFF-test", "audio/wav"))],
)
assert response.status_code == 502
assert calls == [
("POST", "https://api.fish.audio/model"),
("DELETE", "https://api.fish.audio/model/orphan-candidate"),
]
@pytest.mark.parametrize(("text", "expected"), [
("abc", 45),
("漢字", 90),
("🙂", 60),
])
def test_tts_estimate_uses_utf8_bytes(text: str, expected: int):
cost = {"usd": 15 / 1_000_000, "unit": "utf8_byte"}
assert call_resolution._platform_estimate_micro(cost, {}, json.dumps({"text": text}).encode()) == expected