Skip to content

Commit 5f97ffc

Browse files
committed
Fix provider settings model persistence
1 parent 41ed7b2 commit 5f97ffc

10 files changed

Lines changed: 205 additions & 99 deletions

File tree

api/db/schema.sql

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -16,8 +16,9 @@ ALTER TABLE users ADD COLUMN IF NOT EXISTS max_memories_injected INT NOT NULL DE
1616
ALTER TABLE users ADD COLUMN IF NOT EXISTS retrieval_threshold DOUBLE PRECISION NOT NULL DEFAULT 0.5;
1717
ALTER TABLE users ADD COLUMN IF NOT EXISTS dedup_threshold DOUBLE PRECISION NOT NULL DEFAULT 0.95;
1818

19-
ALTER TABLE users ADD COLUMN IF NOT EXISTS extraction_provider TEXT NOT NULL DEFAULT 'openai';
20-
ALTER TABLE users ADD COLUMN IF NOT EXISTS openai_api_key_encrypted BYTEA;
19+
ALTER TABLE users ADD COLUMN IF NOT EXISTS extraction_provider TEXT NOT NULL DEFAULT 'openai';
20+
ALTER TABLE users ADD COLUMN IF NOT EXISTS extraction_model TEXT NOT NULL DEFAULT 'gpt-4o-mini';
21+
ALTER TABLE users ADD COLUMN IF NOT EXISTS openai_api_key_encrypted BYTEA;
2122
ALTER TABLE users ADD COLUMN IF NOT EXISTS gemini_api_key_encrypted BYTEA;
2223
ALTER TABLE users ADD COLUMN IF NOT EXISTS anthropic_api_key_encrypted BYTEA;
2324

api/models/user.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -38,9 +38,10 @@ class UserConfigResponse(BaseModel):
3838
dedup_threshold: float
3939

4040

41-
class UserProviderConfigUpdate(BaseModel):
42-
extraction_provider: str | None = Field(default=None, pattern="^(openai|gemini|ollama|anthropic)$")
43-
openai_api_key: str | None = Field(default=None, max_length=512)
41+
class UserProviderConfigUpdate(BaseModel):
42+
extraction_provider: str | None = Field(default=None, pattern="^(openai|gemini|ollama|anthropic)$")
43+
extraction_model: str | None = Field(default=None, min_length=1, max_length=120)
44+
openai_api_key: str | None = Field(default=None, max_length=512)
4445
gemini_api_key: str | None = Field(default=None, max_length=512)
4546
anthropic_api_key: str | None = Field(default=None, max_length=512)
4647
clear_openai_key: bool = Field(default=False)

api/routes/users.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -142,9 +142,10 @@ async def update_current_user_provider_route(
142142
) -> dict[str, object]:
143143
try:
144144
return await update_user_provider_config(
145-
user["id"],
146-
payload.extraction_provider,
147-
payload.openai_api_key,
145+
user["id"],
146+
payload.extraction_provider,
147+
payload.extraction_model,
148+
payload.openai_api_key,
148149
payload.gemini_api_key,
149150
payload.anthropic_api_key,
150151
payload.clear_openai_key,

api/services/embedding.py

Lines changed: 25 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,19 @@
1-
from sentence_transformers import SentenceTransformer
1+
from typing import Protocol
22

33
from api.config import settings
44

55

6-
_model: SentenceTransformer | None = None
6+
class EmbeddingModel(Protocol):
7+
def encode(self, sentences: str | list[str], normalize_embeddings: bool) -> object:
8+
...
9+
10+
11+
_model: EmbeddingModel | None = None
712

813

914
def load_model() -> None:
15+
from sentence_transformers import SentenceTransformer
16+
1017
global _model
1118
_model = SentenceTransformer(settings.embedding_model)
1219

@@ -15,18 +22,31 @@ def is_model_loaded() -> bool:
1522
return _model is not None
1623

1724

18-
def get_model() -> SentenceTransformer:
25+
def get_model() -> EmbeddingModel:
1926
if _model is None:
2027
raise RuntimeError("Embedding model not loaded")
2128
return _model
2229

2330

2431
def embed(text: str) -> list[float]:
25-
return get_model().encode(text, normalize_embeddings=True).tolist()
32+
return to_float_vector(get_model().encode(text, normalize_embeddings=True))
2633

2734

2835
def embed_batch(texts: list[str]) -> list[list[float]]:
29-
return get_model().encode(texts, normalize_embeddings=True).tolist()
36+
encoded = get_model().encode(texts, normalize_embeddings=True)
37+
if hasattr(encoded, "tolist"):
38+
encoded = encoded.tolist()
39+
if not isinstance(encoded, list):
40+
raise RuntimeError("Embedding batch output is not a list")
41+
return [to_float_vector(vector) for vector in encoded]
42+
43+
44+
def to_float_vector(value: object) -> list[float]:
45+
if hasattr(value, "tolist"):
46+
value = value.tolist()
47+
if not isinstance(value, list):
48+
raise RuntimeError("Embedding output is not a list")
49+
return [float(item) for item in value]
3050

3151

3252
def format_embedding_for_pgvector(embedding: list[float]) -> str:

api/services/proxy.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -61,10 +61,11 @@ async def build_proxy_result(
6161
except Exception as error:
6262
injected_count = 0
6363
logger.warning("Retrieval failed, proceeding without memories: %s", error)
64-
user_row = await db.fetchrow(
65-
"""SELECT id, external_id, extraction_provider,
66-
openai_api_key_encrypted, gemini_api_key_encrypted, anthropic_api_key_encrypted
67-
FROM users WHERE id = $1""",
64+
user_row = await db.fetchrow(
65+
"""SELECT id, external_id, extraction_provider,
66+
extraction_model,
67+
openai_api_key_encrypted, gemini_api_key_encrypted, anthropic_api_key_encrypted
68+
FROM users WHERE id = $1""",
6869
user_id,
6970
)
7071
if user_row is None:
@@ -102,9 +103,10 @@ async def build_proxy_passthrough_result(
102103
from api.config import settings as _settings
103104
fallback_row = {
104105
"id": None,
105-
"external_id": external_id,
106-
"extraction_provider": _settings.extraction_provider,
107-
"openai_api_key_encrypted": None,
106+
"external_id": external_id,
107+
"extraction_provider": _settings.extraction_provider,
108+
"extraction_model": _settings.extraction_model,
109+
"openai_api_key_encrypted": None,
108110
"gemini_api_key_encrypted": None,
109111
"anthropic_api_key_encrypted": None,
110112
}

api/services/users.py

Lines changed: 45 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -3,40 +3,29 @@
33

44
import asyncpg
55

6-
from api.config import settings
7-
from api.services.security import api_key_hashes_match, generate_api_key, hash_api_key
6+
from api.config import settings
7+
from api.services.provider_keys import normalize_provider_name
8+
from api.services.security import api_key_hashes_match, encrypt_provider_key, generate_api_key, hash_api_key
89

9-
from api.services.provider_keys import (
10-
ProviderConfigError,
11-
normalize_provider_name,
12-
)
13-
from api.services.security import encrypt_provider_key
1410

11+
_USER_COLUMNS: str = (
12+
"id, external_id, api_key_hash, created_at, max_memories_injected, "
13+
"retrieval_threshold, dedup_threshold, extraction_provider, "
14+
"extraction_model, openai_api_key_encrypted, gemini_api_key_encrypted, anthropic_api_key_encrypted"
15+
)
1516

16-
_USER_COLUMNS: str = (
17-
"id, external_id, api_key_hash, created_at, max_memories_injected, "
18-
"retrieval_threshold, dedup_threshold, extraction_provider, "
19-
"openai_api_key_encrypted, gemini_api_key_encrypted, anthropic_api_key_encrypted"
20-
)
2117

22-
23-
_PROVIDER_KEY_ENCRYPTED_COLUMNS: dict[str, str] = {
24-
"openai": "openai_api_key_encrypted",
25-
"gemini": "gemini_api_key_encrypted",
26-
"anthropic": "anthropic_api_key_encrypted",
27-
}
28-
29-
30-
class CachedUser(TypedDict):
18+
class CachedUser(TypedDict):
3119
id: object
3220
external_id: str
3321
api_key_hash: str
3422
created_at: object
3523
max_memories_injected: int
3624
retrieval_threshold: float
37-
dedup_threshold: float
38-
extraction_provider: str
39-
cached_at: float
25+
dedup_threshold: float
26+
extraction_provider: str
27+
extraction_model: str
28+
cached_at: float
4029

4130

4231
_user_auth_cache: dict[str, CachedUser] = {}
@@ -111,9 +100,9 @@ async def get_user_by_api_key(api_key: str, db: asyncpg.Connection) -> asyncpg.R
111100
row = await db.fetchrow(
112101
f"""
113102
SELECT users.id, users.external_id, user_api_keys.api_key_hash, users.created_at,
114-
users.max_memories_injected, users.retrieval_threshold, users.dedup_threshold,
115-
users.extraction_provider,
116-
users.openai_api_key_encrypted, users.gemini_api_key_encrypted, users.anthropic_api_key_encrypted
103+
users.max_memories_injected, users.retrieval_threshold, users.dedup_threshold,
104+
users.extraction_provider, users.extraction_model,
105+
users.openai_api_key_encrypted, users.gemini_api_key_encrypted, users.anthropic_api_key_encrypted
117106
FROM user_api_keys
118107
JOIN users ON users.id = user_api_keys.user_id
119108
WHERE user_api_keys.api_key_hash = $1
@@ -248,19 +237,25 @@ async def get_user_provider_config(user_id: object, db: asyncpg.Connection) -> d
248237

249238

250239
async def update_user_provider_config(
251-
user_id: object,
252-
extraction_provider: str | None,
253-
openai_api_key: str | None,
240+
user_id: object,
241+
extraction_provider: str | None,
242+
extraction_model: str | None,
243+
openai_api_key: str | None,
254244
gemini_api_key: str | None,
255245
anthropic_api_key: str | None,
256246
clear_openai_key: bool,
257247
clear_gemini_key: bool,
258248
clear_anthropic_key: bool,
259249
db: asyncpg.Connection,
260250
) -> dict[str, object]:
261-
chosen_provider: str | None = None
262-
if extraction_provider is not None:
263-
chosen_provider = normalize_provider_name(extraction_provider)
251+
chosen_provider: str | None = None
252+
if extraction_provider is not None:
253+
chosen_provider = normalize_provider_name(extraction_provider)
254+
clean_extraction_model: str | None = None
255+
if extraction_model is not None:
256+
clean_extraction_model = extraction_model.strip()
257+
if not clean_extraction_model:
258+
raise ValueError("Extraction model is required")
264259

265260
openai_blob: object = _SENTINEL_NO_UPDATE
266261
gemini_blob: object = _SENTINEL_NO_UPDATE
@@ -280,11 +275,14 @@ async def update_user_provider_config(
280275

281276
assignments: list[str] = []
282277
params: list[object] = []
283-
if chosen_provider is not None:
284-
assignments.append("extraction_provider = $" + str(len(params) + 2))
285-
params.append(chosen_provider)
286-
if openai_blob is not _SENTINEL_NO_UPDATE:
287-
assignments.append("openai_api_key_encrypted = $" + str(len(params) + 2))
278+
if chosen_provider is not None:
279+
assignments.append("extraction_provider = $" + str(len(params) + 2))
280+
params.append(chosen_provider)
281+
if clean_extraction_model is not None:
282+
assignments.append("extraction_model = $" + str(len(params) + 2))
283+
params.append(clean_extraction_model)
284+
if openai_blob is not _SENTINEL_NO_UPDATE:
285+
assignments.append("openai_api_key_encrypted = $" + str(len(params) + 2))
288286
params.append(openai_blob)
289287
if gemini_blob is not _SENTINEL_NO_UPDATE:
290288
assignments.append("gemini_api_key_encrypted = $" + str(len(params) + 2))
@@ -318,19 +316,21 @@ def cache_user_auth(api_key_hash: str, row: asyncpg.Record | dict[str, object])
318316
prune_user_auth_cache()
319317
max_memories_injected = get_row_value(row, "max_memories_injected", settings.max_memories_injected)
320318
retrieval_threshold = get_row_value(row, "retrieval_threshold", settings.retrieval_threshold)
321-
dedup_threshold = get_row_value(row, "dedup_threshold", settings.dedup_threshold)
322-
extraction_provider = get_row_value(row, "extraction_provider", settings.extraction_provider)
323-
_user_auth_cache[api_key_hash] = {
319+
dedup_threshold = get_row_value(row, "dedup_threshold", settings.dedup_threshold)
320+
extraction_provider = get_row_value(row, "extraction_provider", settings.extraction_provider)
321+
extraction_model = get_row_value(row, "extraction_model", settings.extraction_model)
322+
_user_auth_cache[api_key_hash] = {
324323
"id": row["id"],
325324
"external_id": row["external_id"],
326325
"api_key_hash": api_key_hash,
327326
"created_at": row["created_at"],
328327
"max_memories_injected": int(max_memories_injected),
329328
"retrieval_threshold": float(retrieval_threshold),
330-
"dedup_threshold": float(dedup_threshold),
331-
"extraction_provider": str(extraction_provider),
332-
"cached_at": monotonic(),
333-
}
329+
"dedup_threshold": float(dedup_threshold),
330+
"extraction_provider": str(extraction_provider),
331+
"extraction_model": str(extraction_model),
332+
"cached_at": monotonic(),
333+
}
334334
trim_user_auth_cache()
335335

336336

api/test_provider_keys.py

Lines changed: 27 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -59,20 +59,22 @@ def test_mask_provider_key_redacts_middle() -> None:
5959
assert mask_provider_key(None) is None
6060

6161

62-
def test_resolve_user_provider_uses_user_key(monkeypatch) -> None:
63-
monkeypatch.setattr(security.settings, "provider_key_encryption_key", Fernet.generate_key().decode())
64-
user = {
65-
"id": "u1",
66-
"external_id": "alice",
67-
"extraction_provider": "gemini",
68-
"openai_api_key_encrypted": None,
69-
"gemini_api_key_encrypted": encrypt_provider_key("user-gemini"),
70-
"anthropic_api_key_encrypted": None,
71-
}
72-
resolved = provider_keys.resolve_user_provider(user)
73-
assert resolved.name == "gemini"
74-
assert resolved.api_key == "user-gemini"
75-
assert resolved.source == "user"
62+
def test_resolve_user_provider_uses_user_key(monkeypatch) -> None:
63+
monkeypatch.setattr(security.settings, "provider_key_encryption_key", Fernet.generate_key().decode())
64+
user = {
65+
"id": "u1",
66+
"external_id": "alice",
67+
"extraction_provider": "gemini",
68+
"extraction_model": "gemini-1.5-flash",
69+
"openai_api_key_encrypted": None,
70+
"gemini_api_key_encrypted": encrypt_provider_key("user-gemini"),
71+
"anthropic_api_key_encrypted": None,
72+
}
73+
resolved = provider_keys.resolve_user_provider(user)
74+
assert resolved.name == "gemini"
75+
assert resolved.model == "gemini-1.5-flash"
76+
assert resolved.api_key == "user-gemini"
77+
assert resolved.source == "user"
7678

7779

7880
def test_resolve_user_provider_override_key_wins(monkeypatch) -> None:
@@ -129,14 +131,16 @@ def test_settings_accepts_engram_prefixed_env_var(monkeypatch) -> None:
129131

130132
def test_summarize_provider_for_response_masks_key(monkeypatch) -> None:
131133
monkeypatch.setattr(security.settings, "provider_key_encryption_key", Fernet.generate_key().decode())
132-
user = {
133-
"extraction_provider": "openai",
134-
"openai_api_key_encrypted": encrypt_provider_key("sk-supersecretvalue12345"),
135-
"gemini_api_key_encrypted": None,
136-
"anthropic_api_key_encrypted": None,
137-
}
134+
user = {
135+
"extraction_provider": "openai",
136+
"extraction_model": "gpt-4.1-mini",
137+
"openai_api_key_encrypted": encrypt_provider_key("sk-supersecretvalue12345"),
138+
"gemini_api_key_encrypted": None,
139+
"anthropic_api_key_encrypted": None,
140+
}
138141
summary = provider_keys.summarize_provider_for_response(user)
139142
assert summary["has_user_api_key"] is True
140-
assert summary["user_api_key_preview"] is not None
141-
assert "supersecretvalue" not in summary["user_api_key_preview"]
142-
assert summary["extraction_provider"] == "openai"
143+
assert summary["user_api_key_preview"] is not None
144+
assert "supersecretvalue" not in summary["user_api_key_preview"]
145+
assert summary["extraction_provider"] == "openai"
146+
assert summary["extraction_model"] == "gpt-4.1-mini"

0 commit comments

Comments
 (0)