Files
frigate/frigate/test/test_genai_providers.py
T
Nicolas MowenandGitHub 01f3369659
CI / AMD64 Build (push) Canceled after 0s
CI / AMD64 Smoke Test (push) Canceled after 0s
CI / ARM Build (push) Canceled after 0s
CI / Jetson Jetpack 6 (push) Canceled after 0s
CI / AMD64 Extra Build (push) Canceled after 0s
CI / ARM Extra Build (push) Canceled after 0s
CI / Synaptics Build (push) Canceled after 0s
CI / Assemble and push default build (push) Canceled after 0s
Ollama image improvements (#24605)
* Support embeddinggemma for ollama

* Refactor docs

* Allow Ollama to probe cost of image tokens

* Cleanup
2026-10-09 08:01:09 -06:00

1107 lines
40 KiB
Python

"""Smoke tests for GenAI chat providers.
Each provider's ``chat_with_tools_stream`` is driven with a canned "test
response" so the two conversion layers are exercised without any network:
1. Frigate (OpenAI-style) messages -> provider-native request format
2. provider-native response -> Frigate ``("kind", value)`` stream events
These guard against regressions such as tool-call arguments arriving as raw
strings instead of dicts (which crash the ``ToolCall`` model), and multimodal
user content (a list of text/image parts, as injected by ``get_live_context``)
crashing message conversion.
"""
import asyncio
import base64
import json
import unittest
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import requests
from frigate.config import GenAIConfig, GenAIProviderEnum
from frigate.genai import PROVIDERS, load_providers
load_providers()
# A minimal but valid JPEG data URI, mirroring what get_live_context injects.
_TINY_JPEG = base64.b64encode(b"\xff\xd8\xff\xd9").decode("ascii")
_IMAGE_DATA_URI = f"data:image/jpeg;base64,{_TINY_JPEG}"
# Conversation ending in a multimodal user message (text + live image), the
# exact shape the chat endpoint builds after a get_live_context tool result.
MULTIMODAL_MESSAGES = [
{"role": "system", "content": "You are a test assistant."},
{"role": "user", "content": "what do you see on the front camera?"},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "get_live_context",
"arguments": json.dumps({"camera": "front"}),
},
}
],
},
{
"role": "tool",
"tool_call_id": "call_1",
"name": "get_live_context",
"content": json.dumps({"camera": "front"}),
},
{
"role": "user",
"content": [
{
"type": "text",
"text": "Here is the current live image from camera 'front'.",
},
{"type": "image_url", "image_url": {"url": _IMAGE_DATA_URI}},
],
},
]
SIMPLE_MESSAGES = [
{"role": "system", "content": "You are a test assistant."},
{"role": "user", "content": "hello"},
]
TOOLS = [
{
"type": "function",
"function": {
"name": "search_objects",
"description": "Search tracked objects",
"parameters": {
"type": "object",
"properties": {"label": {"type": "string"}},
},
},
}
]
def _make_client(provider: str, **cfg_overrides):
"""Build a provider client offline (no model validation, no network)."""
cfg = GenAIConfig(provider=provider, **cfg_overrides)
cls = PROVIDERS[GenAIProviderEnum(provider)]
return cls(cfg, timeout=5, validate_model=False)
def _collect(client, messages, tools=TOOLS):
"""Drain chat_with_tools_stream into a list of (kind, value) events."""
async def _run():
events = []
async for event in client.chat_with_tools_stream(
messages=messages, tools=tools, tool_choice="auto"
):
events.append(event)
return events
return asyncio.run(_run())
def _final_message(events) -> dict:
messages = [value for (kind, value) in events if kind == "message"]
assert messages, f"stream produced no final message: {events}"
return messages[-1]
def _assert_tool_args_are_dicts(final: dict) -> None:
"""Every returned tool call must expose arguments as a dict, never a string."""
for tool_call in final.get("tool_calls") or []:
assert isinstance(tool_call["arguments"], dict), (
f"tool call arguments must be a dict, got "
f"{type(tool_call['arguments']).__name__}: {tool_call['arguments']!r}"
)
# ---------------------------------------------------------------------------
# OpenAI
# ---------------------------------------------------------------------------
def _openai_tc(index, id=None, name=None, arguments=None):
return SimpleNamespace(
index=index,
id=id,
function=SimpleNamespace(name=name, arguments=arguments),
)
def _openai_chunk(content=None, tool_calls=None, finish_reason=None, usage=None):
delta = SimpleNamespace(
content=content,
tool_calls=tool_calls,
reasoning_content=None,
reasoning=None,
)
choice = SimpleNamespace(delta=delta, finish_reason=finish_reason)
return SimpleNamespace(choices=[choice], usage=usage)
class TestOpenAIProvider(unittest.TestCase):
def _client(self):
return _make_client(
"openai", model="gpt-4o", api_key="k", base_url="http://localhost:9999/v1"
)
def test_stream_tool_call_arguments_are_dict(self):
# Arguments arrive split across chunks, as the real API streams them.
chunks = [
_openai_chunk(
tool_calls=[
_openai_tc(0, id="c1", name="search_objects", arguments='{"label":')
]
),
_openai_chunk(tool_calls=[_openai_tc(0, arguments=' "person"}')]),
_openai_chunk(finish_reason="tool_calls"),
]
client = self._client()
client.provider.chat.completions.create = MagicMock(return_value=iter(chunks))
final = _final_message(_collect(client, SIMPLE_MESSAGES))
self.assertEqual(final["finish_reason"], "tool_calls")
self.assertEqual(len(final["tool_calls"]), 1)
_assert_tool_args_are_dicts(final)
self.assertEqual(final["tool_calls"][0]["arguments"], {"label": "person"})
def test_stream_content_response(self):
chunks = [
_openai_chunk(content="hel"),
_openai_chunk(content="lo"),
_openai_chunk(finish_reason="stop"),
]
client = self._client()
client.provider.chat.completions.create = MagicMock(return_value=iter(chunks))
events = _collect(client, SIMPLE_MESSAGES)
deltas = [v for (k, v) in events if k == "content_delta"]
self.assertEqual("".join(deltas), "hello")
self.assertEqual(_final_message(events)["content"], "hello")
def test_multimodal_message_does_not_crash(self):
client = self._client()
client.provider.chat.completions.create = MagicMock(
return_value=iter([_openai_chunk(content="ok", finish_reason="stop")])
)
# Passing the OpenAI-native multimodal list through must not raise.
final = _final_message(_collect(client, MULTIMODAL_MESSAGES))
self.assertEqual(final["content"], "ok")
# ---------------------------------------------------------------------------
# Gemini
# ---------------------------------------------------------------------------
def _gemini_part(text=None, thought=False, function_call=None, thought_signature=None):
return SimpleNamespace(
text=text,
thought=thought,
function_call=function_call,
thought_signature=thought_signature,
)
def _gemini_chunk(parts, finish_reason=None, usage_metadata=None):
candidate = SimpleNamespace(
content=SimpleNamespace(parts=parts), finish_reason=finish_reason
)
return SimpleNamespace(candidates=[candidate], usage_metadata=usage_metadata)
def _gemini_stream(chunks):
async def _agen(*args, **kwargs):
for chunk in chunks:
yield chunk
return _agen
class TestGeminiProvider(unittest.TestCase):
def _client(self):
return _make_client("gemini", model="gemini-2.5-flash", api_key="k")
def _patch_stream(self, client, chunks):
client.provider = MagicMock()
client.provider.aio.models.generate_content_stream = AsyncMock(
side_effect=_gemini_stream(chunks)
)
def test_stream_parallel_tool_calls_stay_separate_dicts(self):
# Regression: Gemini streams complete function calls. Two calls to the
# same tool must NOT be merged into one concatenated arguments string.
from google.genai.types import FinishReason
chunks = [
_gemini_chunk(
parts=[
_gemini_part(
function_call=SimpleNamespace(
name="search_objects", args={"label": "person"}
)
),
_gemini_part(
function_call=SimpleNamespace(
name="search_objects", args={"limit": 1}
)
),
],
finish_reason=FinishReason.STOP,
),
]
client = self._client()
self._patch_stream(client, chunks)
final = _final_message(_collect(client, SIMPLE_MESSAGES))
self.assertEqual(final["finish_reason"], "tool_calls")
self.assertEqual(len(final["tool_calls"]), 2)
_assert_tool_args_are_dicts(final)
self.assertEqual(final["tool_calls"][0]["arguments"], {"label": "person"})
self.assertEqual(final["tool_calls"][1]["arguments"], {"limit": 1})
def test_stream_content_response(self):
from google.genai.types import FinishReason
chunks = [
_gemini_chunk(parts=[_gemini_part(text="hel")]),
_gemini_chunk(
parts=[_gemini_part(text="lo")], finish_reason=FinishReason.STOP
),
]
client = self._client()
self._patch_stream(client, chunks)
events = _collect(client, SIMPLE_MESSAGES)
deltas = [v for (k, v) in events if k == "content_delta"]
self.assertEqual("".join(deltas), "hello")
self.assertEqual(_final_message(events)["content"], "hello")
def test_multimodal_message_converts_without_crash(self):
# Regression: a user message with list content (text + image_url) used
# to be handed to Part.from_text(text=<list>) and raise ValidationError.
from google.genai.types import FinishReason
client = self._client()
self._patch_stream(
client,
[
_gemini_chunk(
parts=[_gemini_part(text="ok")], finish_reason=FinishReason.STOP
)
],
)
final = _final_message(_collect(client, MULTIMODAL_MESSAGES))
self.assertEqual(final["content"], "ok")
# ---------------------------------------------------------------------------
# Ollama
# ---------------------------------------------------------------------------
class TestOllamaProvider(unittest.TestCase):
def _client(self):
return _make_client("ollama", model="llama3", base_url="http://localhost:9999")
def _run_with_response(self, client, response, messages):
# Ollama uses a non-streaming call when tools are present, via an
# internally-constructed async client.
fake_async = MagicMock()
fake_async.chat = AsyncMock(return_value=response)
with patch(
"frigate.genai.plugins.ollama.OllamaAsyncClient",
return_value=fake_async,
):
return _collect(client, messages)
def test_tool_call_arguments_are_dict(self):
response = {
"message": {
"content": "",
"tool_calls": [
{
"function": {
"name": "search_objects",
"arguments": {"label": "person"},
}
}
],
},
"done": True,
"done_reason": "stop",
"eval_count": 5,
"prompt_eval_count": 3,
"eval_duration": 1_000_000,
}
client = self._client()
final = _final_message(
self._run_with_response(client, response, SIMPLE_MESSAGES)
)
self.assertEqual(final["finish_reason"], "tool_calls")
_assert_tool_args_are_dicts(final)
self.assertEqual(final["tool_calls"][0]["arguments"], {"label": "person"})
def test_multimodal_message_normalizes_image(self):
# Ollama needs content as a string with images pulled into a separate
# field; the normalizer must extract both without crashing.
response = {
"message": {"content": "ok"},
"done": True,
"done_reason": "stop",
}
client = self._client()
final = _final_message(
self._run_with_response(client, response, MULTIMODAL_MESSAGES)
)
self.assertEqual(final["content"], "ok")
def test_normalize_multimodal_content(self):
from frigate.genai.plugins.ollama import _normalize_multimodal_content
text, images = _normalize_multimodal_content(MULTIMODAL_MESSAGES[-1]["content"])
self.assertEqual(
text, "Here is the current live image from camera 'front'.\n[img]"
)
self.assertEqual(images, [b"\xff\xd8\xff\xd9"])
def test_normalize_keeps_text_and_image_order(self):
from frigate.genai.plugins.ollama import _normalize_multimodal_content
text, images = _normalize_multimodal_content(
[
{"type": "text", "text": "intro"},
{"type": "text", "text": "Frame 1"},
{"type": "image_url", "image_url": {"url": _IMAGE_DATA_URI}},
{"type": "text", "text": "Frame 2"},
{"type": "image_url", "image_url": {"url": _IMAGE_DATA_URI}},
]
)
self.assertEqual(text, "intro\nFrame 1\n[img]\nFrame 2\n[img]")
self.assertEqual(len(images), 2)
def test_send_uses_chat_with_captions_before_each_image(self):
client = self._client()
client.provider = MagicMock()
client.provider.chat.return_value = {
"message": {"content": '{"ok": true}'},
"done": True,
"done_reason": "stop",
}
client._supports_thinking_cache = False
result = client._send(
"prompt",
[b"a", b"b"],
{"type": "json_schema", "json_schema": {"schema": {"type": "object"}}},
image_captions=["Frame 1 of 2", "Frame 2 of 2"],
)
self.assertEqual(result, '{"ok": true}')
client.provider.generate.assert_not_called()
params = client.provider.chat.call_args.kwargs
self.assertEqual(
params["messages"],
[
{
"role": "user",
"content": "prompt\nFrame 1 of 2\n[img]\nFrame 2 of 2\n[img]",
"images": [b"a", b"b"],
}
],
)
self.assertEqual(params["format"], {"type": "object"})
self.assertNotIn("think", params)
def test_send_without_captions_puts_images_after_prompt(self):
client = self._client()
client.provider = MagicMock()
client.provider.chat.return_value = {"message": {"content": "ok"}, "done": True}
client._supports_thinking_cache = False
client._send("prompt", [b"a"])
message = client.provider.chat.call_args.kwargs["messages"][0]
self.assertEqual(message["content"], "prompt\n[img]")
self.assertEqual(message["images"], [b"a"])
def test_capabilities_drive_embeddings_and_thinking(self):
client = self._client()
client.provider = MagicMock()
client.provider.show.return_value = {"capabilities": ["embedding", "vision"]}
self.assertTrue(client.supports_embeddings)
self.assertFalse(client.supports_toggleable_thinking)
client.provider.show.assert_called_once()
def test_capability_lookup_failure_is_not_cached(self):
from ollama import ResponseError
client = self._client()
client.provider = MagicMock()
client.provider.show.side_effect = [
ResponseError("unavailable", 503),
{"capabilities": ["embedding"]},
]
self.assertFalse(client.supports_embeddings)
self.assertTrue(client.supports_embeddings)
def test_thinking_rechecked_after_provider_recovers(self):
client = self._client()
client.provider = None
self.assertFalse(client.supports_toggleable_thinking)
client.provider = MagicMock()
client.provider.show.return_value = {"capabilities": ["thinking"]}
self.assertTrue(client.supports_toggleable_thinking)
params = client._build_request_params(
[{"role": "user", "content": "hi"}], None, None, enable_thinking=True
)
self.assertTrue(params["think"])
@staticmethod
def _webp_bytes():
import io
from PIL import Image
buf = io.BytesIO()
Image.new("RGB", (8, 8), (200, 10, 10)).save(buf, format="WEBP")
return buf.getvalue()
def test_embed_posts_text_and_image_items(self):
client = self._client()
client.provider = MagicMock()
client.provider._request_raw.return_value.json.return_value = {
"embeddings": [[0.1] * 768, [0.2] * 768]
}
result = client.embed(texts=["a person"], images=[self._webp_bytes()])
args = client.provider._request_raw.call_args
self.assertEqual(args.args, ("POST", "/api/embed"))
payload = args.kwargs["json"]
self.assertEqual(payload["model"], "llama3")
self.assertEqual(payload["input"][0], "a person")
self.assertEqual(list(payload["input"][1]), ["image"])
# WebP thumbnails are converted to JPEG before being sent
image = base64.b64decode(payload["input"][1]["image"])
self.assertEqual(image[:2], b"\xff\xd8")
self.assertEqual(len(result), 2)
self.assertAlmostEqual(float(result[1][0]), 0.2, places=5)
def test_embed_passes_configured_options(self):
client = _make_client(
"ollama",
model="embeddinggemma-2",
base_url="http://localhost:9999",
provider_options={"options": {"num_ctx": 2048}, "keep_alive": "10m"},
)
client.provider = MagicMock()
client.provider._request_raw.return_value.json.return_value = {
"embeddings": [[0.1] * 768]
}
client.embed(texts=["a"])
payload = client.provider._request_raw.call_args.kwargs["json"]
self.assertEqual(payload["options"], {"num_ctx": 2048})
self.assertEqual(payload["keep_alive"], "10m")
@staticmethod
def _chat_counting_prompt_tokens(**params):
"""Fake chat that reports 10 prompt tokens plus 250 per image."""
images = params["messages"][0].get("images") or []
return {
"message": {"content": "."},
"prompt_eval_count": 10 + 250 * len(images),
}
def test_image_tokens_probed_with_one_token_requests(self):
client = _make_client(
"ollama",
model="qwen3-vl",
base_url="http://localhost:9999",
provider_options={"options": {"num_ctx": 16384}},
)
client.provider = MagicMock()
client.provider.chat.side_effect = self._chat_counting_prompt_tokens
client._supports_thinking_cache = False
self.assertEqual(client.estimate_image_tokens(320, 180), 250)
calls = client.provider.chat.call_args_list
self.assertEqual(len(calls), 2)
for call in calls:
self.assertEqual(
call.kwargs["options"], {"num_ctx": 16384, "num_predict": 1}
)
image = calls[1].kwargs["messages"][0]["images"][0]
self.assertEqual(image[:2], b"\xff\xd8")
def test_image_token_probe_error_uses_heuristic(self):
from ollama import ResponseError
client = self._client()
client.provider = MagicMock()
client.provider.chat.side_effect = [
{"message": {"content": "."}, "prompt_eval_count": 10},
ResponseError("model does not support images", 400),
]
client._supports_thinking_cache = False
self.assertEqual(client.estimate_image_tokens(250, 100), 20)
self.assertEqual(client._image_token_cache, {})
def test_embed_server_error_returns_empty(self):
from ollama import ResponseError
client = self._client()
client.provider = MagicMock()
client.provider._request_raw.side_effect = ResponseError(
"model does not support media embeddings", 400
)
self.assertEqual(client.embed(images=[self._webp_bytes()]), [])
# ---------------------------------------------------------------------------
# llama.cpp
# ---------------------------------------------------------------------------
class _FakeStreamResponse:
def __init__(self, lines):
self._lines = lines
def raise_for_status(self):
return None
async def aiter_lines(self):
for line in self._lines:
yield line
class _FakeStreamCtx:
def __init__(self, lines):
self._resp = _FakeStreamResponse(lines)
async def __aenter__(self):
return self._resp
async def __aexit__(self, *exc):
return False
class _FakeAsyncClient:
def __init__(self, lines):
self._lines = lines
async def __aenter__(self):
return self
async def __aexit__(self, *exc):
return False
def stream(self, method, url, json=None, headers=None):
return _FakeStreamCtx(self._lines)
class TestLlamaCppProvider(unittest.TestCase):
def _client(self):
return _make_client("llamacpp", model="m", base_url="http://localhost:9999")
def _run_with_lines(self, client, lines, messages):
with patch(
"frigate.genai.plugins.llama_cpp.httpx.AsyncClient",
return_value=_FakeAsyncClient(lines),
):
return _collect(client, messages)
def test_stream_tool_call_arguments_are_dict(self):
lines = [
"data: "
+ json.dumps(
{
"choices": [
{
"delta": {
"tool_calls": [
{
"index": 0,
"id": "c1",
"function": {
"name": "search_objects",
"arguments": '{"label":',
},
}
]
}
}
]
}
),
"data: "
+ json.dumps(
{
"choices": [
{
"delta": {
"tool_calls": [
{
"index": 0,
"function": {"arguments": ' "person"}'},
}
]
}
}
]
}
),
"data: "
+ json.dumps({"choices": [{"delta": {}, "finish_reason": "tool_calls"}]}),
"data: [DONE]",
]
client = self._client()
final = _final_message(self._run_with_lines(client, lines, SIMPLE_MESSAGES))
self.assertEqual(final["finish_reason"], "tool_calls")
_assert_tool_args_are_dicts(final)
self.assertEqual(final["tool_calls"][0]["arguments"], {"label": "person"})
def test_stream_content_response(self):
lines = [
"data: " + json.dumps({"choices": [{"delta": {"content": "hel"}}]}),
"data: " + json.dumps({"choices": [{"delta": {"content": "lo"}}]}),
"data: "
+ json.dumps({"choices": [{"delta": {}, "finish_reason": "stop"}]}),
"data: [DONE]",
]
client = self._client()
events = self._run_with_lines(client, lines, SIMPLE_MESSAGES)
deltas = [v for (k, v) in events if k == "content_delta"]
self.assertEqual("".join(deltas), "hello")
self.assertEqual(_final_message(events)["content"], "hello")
def test_multimodal_message_does_not_crash(self):
lines = [
"data: " + json.dumps({"choices": [{"delta": {"content": "ok"}}]}),
"data: "
+ json.dumps({"choices": [{"delta": {}, "finish_reason": "stop"}]}),
"data: [DONE]",
]
client = self._client()
final = _final_message(self._run_with_lines(client, lines, MULTIMODAL_MESSAGES))
self.assertEqual(final["content"], "ok")
def _validated_client(self, server_context_size, provider_options=None):
"""Build a client as if the server reported the given context size."""
cfg = GenAIConfig(
provider="llamacpp",
model="m",
base_url="http://localhost:9999",
provider_options=provider_options or {},
)
info = {
"context_size": server_context_size,
"supports_vision": False,
"supports_audio": False,
"supports_tools": False,
"supports_reasoning": False,
}
cls = PROVIDERS[GenAIProviderEnum.llamacpp]
with patch.object(cls, "_get_model_info", return_value=info):
return cls(cfg, timeout=5)
def test_server_context_size_used_without_override(self):
client = self._validated_client(4096)
self.assertEqual(client.get_context_size(), 4096)
def test_provider_options_context_size_overrides_server(self):
client = self._validated_client(4096, {"context_size": 32768})
self.assertEqual(client.get_context_size(), 32768)
def test_list_models_dedupes_alias_matching_id(self):
client = self._client()
models_data = [
{"id": "qwen3-asr", "aliases": ["qwen3-asr"]},
{"id": "gemma", "aliases": ["gemma", "g4"]},
]
with patch.object(client, "_fetch_models_data", return_value=models_data):
self.assertEqual(client.list_models(), ["g4", "gemma", "qwen3-asr"])
@staticmethod
def _embeddings_response(vectors):
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"object": "list",
"data": [
{"object": "embedding", "index": i, "embedding": v}
for i, v in enumerate(vectors)
],
}
return response
def test_embed_posts_content_arrays_to_v1_embeddings(self):
client = self._client()
response = self._embeddings_response([[0.1] * 768, [0.2] * 768])
with patch.object(client, "_post", return_value=response) as post:
result = client.embed(texts=["a person"], images=[b"not an image"])
url = post.call_args.args[0]
payload = post.call_args.kwargs["json"]
self.assertEqual(url, "http://localhost:9999/v1/embeddings")
self.assertEqual(payload["model"], "m")
self.assertEqual(payload["encoding_format"], "float")
self.assertEqual(
payload["input"][0], {"content": [{"type": "text", "text": "a person"}]}
)
image_parts = payload["input"][1]["content"]
self.assertEqual(image_parts[0]["type"], "image_url")
self.assertEqual(
image_parts[0]["image_url"]["url"],
"data:image/jpeg;base64," + base64.b64encode(b"not an image").decode(),
)
self.assertEqual(image_parts[1], {"type": "text", "text": "\n"})
self.assertEqual(len(result), 2)
self.assertAlmostEqual(float(result[1][0]), 0.2, places=5)
def test_embed_normalizes_dimension(self):
client = self._client()
response = self._embeddings_response([[1.0] * 1024, [1.0] * 512])
with patch.object(client, "_post", return_value=response):
result = client.embed(texts=["long", "short"])
self.assertEqual([r.shape for r in result], [(768,), (768,)])
self.assertEqual(float(result[1][-1]), 0.0)
@staticmethod
def _post_counting_prompt_tokens(url, json=None, timeout=None):
"""Fake chat completion: 12 prompt tokens plus 300 per image part."""
content = json["messages"][0]["content"]
images = [p for p in content if p["type"] == "image_url"]
response = MagicMock()
response.json.return_value = {
"usage": {"prompt_tokens": 12 + 300 * len(images)}
}
return response
def test_image_tokens_probed_once_per_dimension(self):
client = self._client()
with patch.object(
client, "_post", side_effect=self._post_counting_prompt_tokens
) as post:
self.assertEqual(client.estimate_image_tokens(320, 180), 300)
self.assertEqual(client.estimate_image_tokens(320, 180), 300)
self.assertEqual(client.estimate_image_tokens(640, 360), 300)
# one shared text baseline, then one image request per new dimension
self.assertEqual(post.call_count, 3)
payload = post.call_args_list[0].kwargs["json"]
self.assertEqual(payload["max_tokens"], 1)
self.assertEqual(
post.call_args_list[0].args[0], "http://localhost:9999/v1/chat/completions"
)
def test_image_token_probe_failure_is_not_cached(self):
client = self._client()
with patch.object(
client, "_post", side_effect=requests.exceptions.ConnectionError("down")
):
self.assertEqual(client.estimate_image_tokens(250, 100), 20)
self.assertEqual(client._image_token_cache, {})
self.assertIsNone(client._text_baseline_tokens)
def test_embed_request_error_returns_empty(self):
client = self._client()
response = MagicMock()
response.raise_for_status.side_effect = requests.exceptions.HTTPError("400")
with patch.object(client, "_post", return_value=response):
self.assertEqual(client.embed(texts=["a"]), [])
# ---------------------------------------------------------------------------
# transcribe role
# ---------------------------------------------------------------------------
WAV_BYTES = b"RIFF$\x00\x00\x00WAVEfmt "
class TestOpenAITranscribe(unittest.TestCase):
def _client(self):
return _make_client(
"openai",
model="gpt-4o-transcribe",
api_key="k",
base_url="http://localhost:9999/v1",
runtime_options={"temperature": 0.7},
)
def test_supports_transcription(self):
self.assertTrue(self._client().supports_transcription)
def test_passes_file_tuple_and_language(self):
client = self._client()
create = MagicMock(return_value=" hello there ")
client.provider = SimpleNamespace(
audio=SimpleNamespace(transcriptions=SimpleNamespace(create=create))
)
self.assertEqual(client.transcribe(WAV_BYTES, language="en"), "hello there")
kwargs = create.call_args.kwargs
self.assertEqual(kwargs["model"], "gpt-4o-transcribe")
self.assertEqual(kwargs["file"], ("audio.wav", WAV_BYTES, "audio/wav"))
self.assertEqual(kwargs["language"], "en")
self.assertEqual(kwargs["response_format"], "text")
def test_does_not_forward_runtime_options(self):
"""runtime_options are chat parameters; /audio/transcriptions rejects them."""
client = self._client()
create = MagicMock(return_value="hi")
client.provider = SimpleNamespace(
audio=SimpleNamespace(transcriptions=SimpleNamespace(create=create))
)
client.transcribe(WAV_BYTES)
self.assertNotIn("temperature", create.call_args.kwargs)
def test_gpt_transcribe_uses_languages_array(self):
"""gpt-transcribe replaced `language` with a `languages` array."""
client = _make_client("openai", model="gpt-transcribe", api_key="k")
create = MagicMock(return_value="hi")
client.provider = SimpleNamespace(
audio=SimpleNamespace(transcriptions=SimpleNamespace(create=create))
)
client.transcribe(WAV_BYTES, language="en")
kwargs = create.call_args.kwargs
self.assertEqual(kwargs["extra_body"], {"languages": ["en"]})
# sending both fields is rejected by the API
self.assertNotIn("language", kwargs)
def test_older_models_use_singular_language(self):
for model in ("gpt-4o-transcribe", "whisper-1"):
with self.subTest(model):
client = _make_client("openai", model=model, api_key="k")
create = MagicMock(return_value="hi")
client.provider = SimpleNamespace(
audio=SimpleNamespace(transcriptions=SimpleNamespace(create=create))
)
client.transcribe(WAV_BYTES, language="en")
kwargs = create.call_args.kwargs
self.assertEqual(kwargs["language"], "en")
self.assertNotIn("extra_body", kwargs)
def test_object_response_form(self):
client = self._client()
create = MagicMock(return_value=SimpleNamespace(text="hi"))
client.provider = SimpleNamespace(
audio=SimpleNamespace(transcriptions=SimpleNamespace(create=create))
)
self.assertEqual(client.transcribe(WAV_BYTES), "hi")
def test_error_returns_none(self):
client = self._client()
create = MagicMock(side_effect=RuntimeError("boom"))
client.provider = SimpleNamespace(
audio=SimpleNamespace(transcriptions=SimpleNamespace(create=create))
)
self.assertIsNone(client.transcribe(WAV_BYTES))
class TestAzureOpenAITranscribe(unittest.TestCase):
def _client(self):
return _make_client(
"azure_openai",
model="my-deployment",
api_key="k",
base_url="https://example.openai.azure.com/?api-version=2024-06-01",
)
def test_routes_through_azure_client(self):
from openai import AzureOpenAI
client = self._client()
self.assertIsInstance(client.provider, AzureOpenAI)
self.assertTrue(client.supports_transcription)
def test_transcribe_inherited(self):
client = self._client()
create = MagicMock(return_value="azure text")
client.provider = SimpleNamespace(
audio=SimpleNamespace(transcriptions=SimpleNamespace(create=create))
)
self.assertEqual(client.transcribe(WAV_BYTES, language="fr"), "azure text")
self.assertEqual(create.call_args.kwargs["model"], "my-deployment")
class TestGeminiTranscribe(unittest.TestCase):
def _client(self):
return _make_client("gemini", model="gemini-2.0-flash", api_key="k")
def test_supports_transcription(self):
self.assertTrue(self._client().supports_transcription)
def test_sends_audio_part(self):
client = self._client()
generate = MagicMock(return_value=SimpleNamespace(text=" spoken words "))
client.provider = SimpleNamespace(
models=SimpleNamespace(generate_content=generate)
)
self.assertEqual(client.transcribe(WAV_BYTES, language="en"), "spoken words")
contents = generate.call_args.kwargs["contents"]
audio_parts = [
p for p in contents if getattr(p, "inline_data", None) is not None
]
self.assertEqual(len(audio_parts), 1)
self.assertEqual(audio_parts[0].inline_data.mime_type, "audio/wav")
self.assertEqual(audio_parts[0].inline_data.data, WAV_BYTES)
def test_oversized_payload_is_skipped(self):
from frigate.genai.plugins.gemini import GEMINI_MAX_INLINE_BYTES
client = self._client()
generate = MagicMock()
client.provider = SimpleNamespace(
models=SimpleNamespace(generate_content=generate)
)
self.assertIsNone(client.transcribe(b"\x00" * (GEMINI_MAX_INLINE_BYTES + 1)))
generate.assert_not_called()
class TestLlamaCppTranscribe(unittest.TestCase):
def _client(self, supports_audio: bool):
cfg = GenAIConfig(
provider="llamacpp",
model="m",
base_url="http://localhost:9999",
)
info = {
"context_size": 4096,
"supports_vision": False,
"supports_audio": supports_audio,
"supports_tools": False,
"supports_reasoning": False,
}
cls = PROVIDERS[GenAIProviderEnum.llamacpp]
with patch.object(cls, "_get_model_info", return_value=info):
return cls(cfg, timeout=5)
def test_supports_transcription_tracks_supports_audio(self):
self.assertTrue(self._client(True).supports_transcription)
self.assertFalse(self._client(False).supports_transcription)
@staticmethod
def _transcriptions_response(text: str = " transcript "):
response = MagicMock()
response.status_code = 200
response.json.return_value = {"text": text}
return response
@staticmethod
def _chat_response(content: str = " fallback transcript "):
response = MagicMock()
response.status_code = 200
response.json.return_value = {"choices": [{"message": {"content": content}}]}
return response
def test_posts_multipart_to_transcriptions(self):
client = self._client(True)
with patch.object(
client, "_post", return_value=self._transcriptions_response()
) as post:
self.assertEqual(client.transcribe(WAV_BYTES, language="en"), "transcript")
self.assertTrue(post.call_args.args[0].endswith("/v1/audio/transcriptions"))
self.assertEqual(
post.call_args.kwargs["files"]["file"],
("audio.wav", WAV_BYTES, "audio/wav"),
)
self.assertEqual(post.call_args.kwargs["data"]["language"], "en")
def test_omits_language_when_not_set(self):
"""An unset language is what lets the model detect one itself."""
client = self._client(True)
with patch.object(
client, "_post", return_value=self._transcriptions_response()
) as post:
client.transcribe(WAV_BYTES)
self.assertNotIn("language", post.call_args.kwargs["data"])
def test_falls_back_to_chat_completions_on_404(self):
"""Servers predating llama.cpp#21863 have no transcriptions route."""
client = self._client(True)
missing = MagicMock()
missing.status_code = 404
with patch.object(
client, "_post", side_effect=[missing, self._chat_response()]
) as post:
self.assertEqual(
client.transcribe(WAV_BYTES, language="en"), "fallback transcript"
)
urls = [call.args[0] for call in post.call_args_list]
self.assertTrue(urls[0].endswith("/v1/audio/transcriptions"))
self.assertTrue(urls[1].endswith("/v1/chat/completions"))
payload = post.call_args_list[1].kwargs["json"]
content = payload["messages"][0]["content"]
audio_parts = [p for p in content if p["type"] == "input_audio"]
self.assertEqual(len(audio_parts), 1)
self.assertEqual(audio_parts[0]["input_audio"]["format"], "wav")
self.assertEqual(
base64.b64decode(audio_parts[0]["input_audio"]["data"]), WAV_BYTES
)
def test_audio_unsupported_returns_none(self):
client = self._client(False)
with patch.object(client, "_post") as post:
self.assertIsNone(client.transcribe(WAV_BYTES))
post.assert_not_called()
class TestImageTokenEstimate(unittest.TestCase):
def test_provider_without_token_counts_uses_heuristic(self):
client = _make_client("gemini", model="m", api_key="k")
self.assertEqual(client.estimate_image_tokens(250, 100), 20)
self.assertEqual(client._image_token_cache, {})
class TestBaseClientTranscribe(unittest.TestCase):
"""Providers that don't implement the role must be inert, not broken."""
def test_ollama_reports_and_returns_nothing(self):
client = _make_client("ollama", model="llava", base_url="http://localhost:9999")
self.assertFalse(client.supports_transcription)
self.assertIsNone(client.transcribe(WAV_BYTES, language="en"))
if __name__ == "__main__":
unittest.main()