-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmemory_embeddings.py
More file actions
78 lines (63 loc) · 2.25 KB
/
Copy pathmemory_embeddings.py
File metadata and controls
78 lines (63 loc) · 2.25 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
from __future__ import annotations
"""
Gemini embeddings client for Azurro Memory Vault (v1 API).
Uses models/gemini-embedding-2-preview and truncates to 768 dimensions
for storage in pgvector.
"""
import os
from typing import List
from google import genai
_API_KEY = os.getenv("GEMINI_API_KEY", "").strip()
_EMBED_MODEL = os.getenv("AZURRO_EMBED_MODEL", "models/gemini-embedding-2-preview")
_DIM = 768
_client: genai.Client | None = None
def _get_client() -> genai.Client:
global _client
if not _API_KEY:
raise RuntimeError("GEMINI_API_KEY is not set for embeddings")
if _client is None:
_client = genai.Client(api_key=_API_KEY)
return _client
def _truncate(values: List[float]) -> List[float]:
if len(values) >= _DIM:
return values[:_DIM]
# pad if shorter (unlikely)
return values + [0.0] * (_DIM - len(values))
def embed_text(text: str) -> List[float]:
"""
Return a 768-dim embedding vector for the given text using Gemini embeddings.
"""
client = _get_client()
resp = client.models.embed_content(
model=_EMBED_MODEL,
contents=text,
)
emb = getattr(resp, "embedding", None)
if emb and getattr(emb, "values", None):
return _truncate(list(emb.values))
embs = getattr(resp, "embeddings", None)
if embs and getattr(embs[0], "values", None):
return _truncate(list(embs[0].values))
raise RuntimeError("No embedding returned from Gemini")
def embed_texts(texts: List[str]) -> List[List[float]]:
"""
Batch embed a list of texts. Best-effort: falls back to per-item on error.
"""
if not texts:
return []
client = _get_client()
try:
resp = client.models.embed_content(
model=_EMBED_MODEL,
contents=texts,
)
embs = getattr(resp, "embeddings", None)
if not embs:
# Fallback to single embedding if API collapsed
emb = getattr(resp, "embedding", None)
if emb and getattr(emb, "values", None):
return [_truncate(list(emb.values))]
raise RuntimeError("No embeddings returned from Gemini (batch)")
return [_truncate(list(e.values)) for e in embs]
except Exception:
return [embed_text(t) for t in texts]