mirror of
https://github.com/blakeblackshear/frigate.git
synced 2026-08-12 13:51:12 +03:00
fix(genai): drop twelvelabs SDK dep, call Marengo via REST
This commit is contained in:
@@ -50,7 +50,6 @@ transformers == 4.45.*
|
|||||||
google-genai == 1.58.*
|
google-genai == 1.58.*
|
||||||
ollama == 0.6.*
|
ollama == 0.6.*
|
||||||
openai == 1.65.*
|
openai == 1.65.*
|
||||||
twelvelabs == 1.2.*
|
|
||||||
# push notifications
|
# push notifications
|
||||||
py-vapid == 1.9.*
|
py-vapid == 1.9.*
|
||||||
pywebpush == 2.0.*
|
pywebpush == 2.0.*
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ upstream and is consistent for both text and image inputs.
|
|||||||
import logging
|
import logging
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import requests
|
||||||
|
|
||||||
from frigate.config import GenAIProviderEnum
|
from frigate.config import GenAIProviderEnum
|
||||||
from frigate.genai import GenAIClient, register_genai_provider
|
from frigate.genai import GenAIClient, register_genai_provider
|
||||||
@@ -29,26 +30,28 @@ logger = logging.getLogger(__name__)
|
|||||||
# Default Marengo model. Overridable via the `model` config field.
|
# Default Marengo model. Overridable via the `model` config field.
|
||||||
DEFAULT_MODEL = "marengo3.0"
|
DEFAULT_MODEL = "marengo3.0"
|
||||||
|
|
||||||
|
# Marengo embed REST endpoint. No SDK is needed — this is a plain multipart POST
|
||||||
|
# made through Frigate's existing `requests` dependency.
|
||||||
|
EMBED_URL = "https://api.twelvelabs.io/v1.3/embed"
|
||||||
|
|
||||||
|
|
||||||
@register_genai_provider(GenAIProviderEnum.twelvelabs)
|
@register_genai_provider(GenAIProviderEnum.twelvelabs)
|
||||||
class TwelveLabsClient(GenAIClient):
|
class TwelveLabsClient(GenAIClient):
|
||||||
"""GenAI client for Frigate using TwelveLabs Marengo embeddings."""
|
"""GenAI client for Frigate using TwelveLabs Marengo embeddings."""
|
||||||
|
|
||||||
def _init_provider(self):
|
def _init_provider(self):
|
||||||
"""Initialize the TwelveLabs SDK client."""
|
"""Validate config for the TwelveLabs REST provider.
|
||||||
try:
|
|
||||||
from twelvelabs import TwelveLabs
|
|
||||||
except ImportError:
|
|
||||||
logger.error(
|
|
||||||
"The twelvelabs package is required for the TwelveLabs provider."
|
|
||||||
)
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
The provider is just an HTTPS API, so there is no client object to
|
||||||
|
build — the API key is the only thing required. A non-None sentinel is
|
||||||
|
returned so the shared ``ensure_provider``/initialization machinery
|
||||||
|
treats the provider as available.
|
||||||
|
"""
|
||||||
if not self.genai_config.api_key:
|
if not self.genai_config.api_key:
|
||||||
logger.error("TwelveLabs provider requires an api_key.")
|
logger.error("TwelveLabs provider requires an api_key.")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
return TwelveLabs(api_key=self.genai_config.api_key)
|
return self.genai_config.api_key
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def _model(self) -> str:
|
def _model(self) -> str:
|
||||||
@@ -69,6 +72,8 @@ class TwelveLabsClient(GenAIClient):
|
|||||||
sent one at a time. Returns one 512-dim float32 vector per input, in
|
sent one at a time. Returns one 512-dim float32 vector per input, in
|
||||||
order (texts first, then images). The shared GenAIEmbedding adapter
|
order (texts first, then images). The shared GenAIEmbedding adapter
|
||||||
pads these to Frigate's 768-dim search schema.
|
pads these to Frigate's 768-dim search schema.
|
||||||
|
|
||||||
|
Calls the Marengo REST endpoint directly via ``requests`` — no SDK.
|
||||||
"""
|
"""
|
||||||
if self.provider is None:
|
if self.provider is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -93,28 +98,42 @@ class TwelveLabsClient(GenAIClient):
|
|||||||
def _embed_one(
|
def _embed_one(
|
||||||
self, text: str | None = None, image: bytes | None = None
|
self, text: str | None = None, image: bytes | None = None
|
||||||
) -> np.ndarray | None:
|
) -> np.ndarray | None:
|
||||||
"""Embed a single text or image input, returning a float32 vector."""
|
"""Embed a single text or image input, returning a float32 vector.
|
||||||
try:
|
|
||||||
if text is not None:
|
|
||||||
response = self.provider.embed.create(
|
|
||||||
model_name=self._model,
|
|
||||||
text=text,
|
|
||||||
request_options={"timeout_in_seconds": self.timeout},
|
|
||||||
)
|
|
||||||
result = response.text_embedding
|
|
||||||
else:
|
|
||||||
response = self.provider.embed.create(
|
|
||||||
model_name=self._model,
|
|
||||||
image_file=image,
|
|
||||||
request_options={"timeout_in_seconds": self.timeout},
|
|
||||||
)
|
|
||||||
result = response.image_embedding
|
|
||||||
|
|
||||||
if result is None or not result.segments:
|
Posts a multipart form to the Marengo embed endpoint (``model_name`` plus
|
||||||
|
either a ``text`` or an ``image_file`` part). The endpoint requires
|
||||||
|
multipart/form-data, so every field — including text — is passed via
|
||||||
|
``files`` (the ``(None, value)`` form makes requests emit a multipart
|
||||||
|
text part). ``self.provider`` holds the validated API key. The 512-dim
|
||||||
|
vector is at ``<text|image>_embedding.segments[0].float`` in the JSON
|
||||||
|
response.
|
||||||
|
"""
|
||||||
|
headers = {"x-api-key": self.provider}
|
||||||
|
files: dict = {"model_name": (None, self._model)}
|
||||||
|
|
||||||
|
if text is not None:
|
||||||
|
files["text"] = (None, text)
|
||||||
|
result_key = "text_embedding"
|
||||||
|
else:
|
||||||
|
files["image_file"] = ("image.jpg", image, "image/jpeg")
|
||||||
|
result_key = "image_embedding"
|
||||||
|
|
||||||
|
try:
|
||||||
|
response = requests.post(
|
||||||
|
EMBED_URL,
|
||||||
|
headers=headers,
|
||||||
|
files=files,
|
||||||
|
timeout=self.timeout,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
result = response.json().get(result_key) or {}
|
||||||
|
segments = result.get("segments") or []
|
||||||
|
|
||||||
|
if not segments:
|
||||||
logger.warning("TwelveLabs returned no embedding for input.")
|
logger.warning("TwelveLabs returned no embedding for input.")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
return np.array(result.segments[0].float_, dtype=np.float32)
|
return np.array(segments[0]["float"], dtype=np.float32)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("TwelveLabs returned an error: %s", e)
|
logger.warning("TwelveLabs returned an error: %s", e)
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
import io
|
import io
|
||||||
import os
|
import os
|
||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
@@ -24,91 +24,85 @@ def _make_config(model: str = "") -> GenAIConfig:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _segment(values):
|
def _response(key: str, values):
|
||||||
"""Mimic the SDK BaseSegment shape (a `float_` list per segment)."""
|
"""Mimic the Marengo REST JSON: ``{<key>: {segments: [{float: [...]}]}}``."""
|
||||||
seg = MagicMock()
|
resp = MagicMock()
|
||||||
seg.float_ = values
|
resp.raise_for_status.return_value = None
|
||||||
return seg
|
resp.json.return_value = {key: {"segments": [{"float": values}]}}
|
||||||
|
return resp
|
||||||
|
|
||||||
|
|
||||||
class TestTwelveLabsEmbedNoNetwork(unittest.TestCase):
|
class TestTwelveLabsEmbedNoNetwork(unittest.TestCase):
|
||||||
"""Unit tests with the SDK client mocked — no network access."""
|
"""Unit tests with ``requests`` mocked — no network access, no SDK."""
|
||||||
|
|
||||||
def _client_with_provider(self, provider) -> TwelveLabsClient:
|
def _client(self) -> TwelveLabsClient:
|
||||||
client = TwelveLabsClient.__new__(TwelveLabsClient)
|
client = TwelveLabsClient.__new__(TwelveLabsClient)
|
||||||
client.genai_config = _make_config()
|
client.genai_config = _make_config()
|
||||||
client.timeout = 120
|
client.timeout = 120
|
||||||
client.provider = provider
|
client.provider = "test-key"
|
||||||
return client
|
return client
|
||||||
|
|
||||||
def test_text_embedding_returns_vector(self):
|
@patch("frigate.genai.plugins.twelvelabs.requests.post")
|
||||||
provider = MagicMock()
|
def test_text_embedding_returns_vector(self, post):
|
||||||
response = MagicMock()
|
post.return_value = _response("text_embedding", [0.1, 0.2, 0.3])
|
||||||
response.text_embedding.segments = [_segment([0.1, 0.2, 0.3])]
|
|
||||||
provider.embed.create.return_value = response
|
|
||||||
|
|
||||||
client = self._client_with_provider(provider)
|
out = self._client().embed(texts=["a person walking a dog"])
|
||||||
out = client.embed(texts=["a person walking a dog"])
|
|
||||||
|
|
||||||
self.assertEqual(len(out), 1)
|
self.assertEqual(len(out), 1)
|
||||||
self.assertIsInstance(out[0], np.ndarray)
|
self.assertIsInstance(out[0], np.ndarray)
|
||||||
self.assertEqual(out[0].dtype, np.float32)
|
self.assertEqual(out[0].dtype, np.float32)
|
||||||
np.testing.assert_allclose(out[0], [0.1, 0.2, 0.3], rtol=1e-6)
|
np.testing.assert_allclose(out[0], [0.1, 0.2, 0.3], rtol=1e-6)
|
||||||
|
|
||||||
_, kwargs = provider.embed.create.call_args
|
_, kwargs = post.call_args
|
||||||
self.assertEqual(kwargs["model_name"], DEFAULT_MODEL)
|
self.assertEqual(kwargs["files"]["model_name"][1], DEFAULT_MODEL)
|
||||||
self.assertEqual(kwargs["text"], "a person walking a dog")
|
self.assertEqual(kwargs["files"]["text"][1], "a person walking a dog")
|
||||||
|
self.assertEqual(kwargs["headers"]["x-api-key"], "test-key")
|
||||||
|
self.assertNotIn("image_file", kwargs["files"])
|
||||||
|
|
||||||
def test_image_embedding_uses_image_file(self):
|
@patch("frigate.genai.plugins.twelvelabs.requests.post")
|
||||||
provider = MagicMock()
|
def test_image_embedding_uses_image_file(self, post):
|
||||||
response = MagicMock()
|
post.return_value = _response("image_embedding", [1.0, 2.0])
|
||||||
response.image_embedding.segments = [_segment([1.0, 2.0])]
|
|
||||||
provider.embed.create.return_value = response
|
|
||||||
|
|
||||||
client = self._client_with_provider(provider)
|
out = self._client().embed(images=[b"\xff\xd8\xff jpeg bytes"])
|
||||||
out = client.embed(images=[b"\xff\xd8\xff jpeg bytes"])
|
|
||||||
|
|
||||||
self.assertEqual(len(out), 1)
|
self.assertEqual(len(out), 1)
|
||||||
_, kwargs = provider.embed.create.call_args
|
_, kwargs = post.call_args
|
||||||
self.assertEqual(kwargs["image_file"], b"\xff\xd8\xff jpeg bytes")
|
self.assertEqual(kwargs["files"]["image_file"][1], b"\xff\xd8\xff jpeg bytes")
|
||||||
self.assertNotIn("text", kwargs)
|
self.assertNotIn("text", kwargs["files"])
|
||||||
|
|
||||||
def test_custom_model_name_is_used(self):
|
@patch("frigate.genai.plugins.twelvelabs.requests.post")
|
||||||
provider = MagicMock()
|
def test_custom_model_name_is_used(self, post):
|
||||||
response = MagicMock()
|
post.return_value = _response("text_embedding", [0.0])
|
||||||
response.text_embedding.segments = [_segment([0.0])]
|
|
||||||
provider.embed.create.return_value = response
|
|
||||||
|
|
||||||
client = self._client_with_provider(provider)
|
client = self._client()
|
||||||
client.genai_config = _make_config(model="marengo-custom")
|
client.genai_config = _make_config(model="marengo-custom")
|
||||||
client.embed(texts=["x"])
|
client.embed(texts=["x"])
|
||||||
|
|
||||||
_, kwargs = provider.embed.create.call_args
|
_, kwargs = post.call_args
|
||||||
self.assertEqual(kwargs["model_name"], "marengo-custom")
|
self.assertEqual(kwargs["files"]["model_name"][1], "marengo-custom")
|
||||||
|
|
||||||
def test_empty_segments_are_skipped(self):
|
@patch("frigate.genai.plugins.twelvelabs.requests.post")
|
||||||
provider = MagicMock()
|
def test_empty_segments_are_skipped(self, post):
|
||||||
response = MagicMock()
|
resp = MagicMock()
|
||||||
response.text_embedding = None
|
resp.raise_for_status.return_value = None
|
||||||
provider.embed.create.return_value = response
|
resp.json.return_value = {"text_embedding": {"segments": []}}
|
||||||
|
post.return_value = resp
|
||||||
|
|
||||||
client = self._client_with_provider(provider)
|
self.assertEqual(self._client().embed(texts=["x"]), [])
|
||||||
self.assertEqual(client.embed(texts=["x"]), [])
|
|
||||||
|
|
||||||
def test_api_error_is_swallowed(self):
|
@patch("frigate.genai.plugins.twelvelabs.requests.post")
|
||||||
provider = MagicMock()
|
def test_api_error_is_swallowed(self, post):
|
||||||
provider.embed.create.side_effect = RuntimeError("boom")
|
post.side_effect = RuntimeError("boom")
|
||||||
|
|
||||||
client = self._client_with_provider(provider)
|
self.assertEqual(self._client().embed(texts=["x"]), [])
|
||||||
self.assertEqual(client.embed(texts=["x"]), [])
|
|
||||||
|
|
||||||
def test_no_provider_returns_empty(self):
|
def test_no_provider_returns_empty(self):
|
||||||
client = self._client_with_provider(None)
|
client = self._client()
|
||||||
|
client.provider = None
|
||||||
self.assertEqual(client.embed(texts=["x"]), [])
|
self.assertEqual(client.embed(texts=["x"]), [])
|
||||||
|
|
||||||
def test_no_inputs_returns_empty(self):
|
def test_no_inputs_returns_empty(self):
|
||||||
client = self._client_with_provider(MagicMock())
|
self.assertEqual(self._client().embed(), [])
|
||||||
self.assertEqual(client.embed(), [])
|
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipUnless(
|
@unittest.skipUnless(
|
||||||
|
|||||||
Reference in New Issue
Block a user