Skip to content

Commit 41fb6a3

Browse files
lesebclaude
andcommitted
feat(vector-io): enforce metadata_store when access control is enabled
Gate vector store providers to require metadata_store (AuthorizedSqlStore) when access control policies are active, ensuring tenant isolation cannot be silently bypassed. Auto-inject metadata_store from server-level storage.stores.vector_stores into all vector_io providers via the resolver, so operators configure it once instead of per-provider. Fix batch cleanup to use unfiltered SQL access for internal bookkeeping. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com> Signed-off-by: Sébastien Han <seb@redhat.com>
1 parent 7381eb6 commit 41fb6a3

5 files changed

Lines changed: 250 additions & 12 deletions

File tree

.github/workflows/integration-responses-conversations-auth-tests.yml

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -167,13 +167,18 @@ jobs:
167167
prompts:
168168
table_name: prompts
169169
backend: sql_default
170+
vector_stores:
171+
table_name: vector_store_metadata
172+
backend: sql_default
170173
models:
171174
- model_id: openai/gpt-4o
172175
model_type: llm
173176
provider_id: openai
174177
- model_id: sentence-transformers/nomic-ai/nomic-embed-text-v1.5
175178
model_type: embedding
176179
provider_id: sentence-transformers
180+
metadata:
181+
embedding_dimension: 768
177182
server:
178183
port: 8321
179184
auth:

src/ogx/core/resolver.py

Lines changed: 21 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -400,10 +400,29 @@ async def instantiate_provider(
400400

401401
logger.debug("Instantiating provider", provider_id=provider.provider_id, module=provider_spec.module)
402402
module = importlib.import_module(provider_spec.module)
403+
404+
def _inject_config_defaults(config_type: type[Any], provider_config: dict[str, Any]) -> dict[str, Any]:
405+
fields = getattr(config_type, "__fields__", None)
406+
if fields is None:
407+
return provider_config
408+
409+
# Inject vector_stores_config for providers that need it (introspection-based).
410+
# Only inject if vector_stores is provided, otherwise let default_factory handle it.
411+
if "vector_stores_config" in fields and run_config.vector_stores is not None:
412+
provider_config["vector_stores_config"] = run_config.vector_stores
413+
414+
# Inject metadata_store from server stores config when not explicitly configured.
415+
if "metadata_store" in fields:
416+
if provider_config.get("metadata_store") is None and run_config.storage.stores.vector_stores is not None:
417+
provider_config["metadata_store"] = run_config.storage.stores.vector_stores.model_dump()
418+
419+
return provider_config
420+
403421
args = []
404422
if isinstance(provider_spec, RemoteProviderSpec):
405423
config_type = instantiate_class_type(provider_spec.config_class)
406-
config = config_type(**provider.config)
424+
provider_config = _inject_config_defaults(config_type, provider.config.copy())
425+
config = config_type(**provider_config)
407426

408427
method = "get_adapter_impl"
409428
args = [config, deps]
@@ -423,15 +442,8 @@ async def instantiate_provider(
423442
args = [provider_spec.api, inner_impls, deps, dist_registry, policy]
424443
else:
425444
method = "get_provider_impl"
426-
provider_config = provider.config.copy()
427-
428-
# Inject vector_stores_config for providers that need it (introspection-based)
429445
config_type = instantiate_class_type(provider_spec.config_class)
430-
if hasattr(config_type, "__fields__") and "vector_stores_config" in config_type.__fields__:
431-
# Only inject if vector_stores is provided, otherwise let default_factory handle it
432-
if run_config.vector_stores is not None:
433-
provider_config["vector_stores_config"] = run_config.vector_stores
434-
446+
provider_config = _inject_config_defaults(config_type, provider.config.copy())
435447
config = config_type(**provider_config)
436448
args = [config, deps]
437449
if "policy" in inspect.signature(getattr(module, method)).parameters:

src/ogx/providers/utils/memory/openai_vector_store_mixin.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -599,15 +599,18 @@ async def _delete_openai_vector_store_file_batch(self, batch_id: str) -> None:
599599
async def _cleanup_expired_file_batches(self) -> None:
600600
"""Clean up expired file batches from persistent storage."""
601601
if self.metadata_store:
602-
results = await self.metadata_store.fetch_all(table=TABLE_VECTOR_STORE_FILE_BATCHES)
602+
rows = await self._fetch_all_metadata_rows_unfiltered(table=TABLE_VECTOR_STORE_FILE_BATCHES)
603603
current_time = int(time.time())
604604
expired_count = 0
605-
for row in results.data:
605+
for row in rows:
606606
info = row["batch_data"]
607607
expires_at = info.get("expires_at")
608608
if expires_at and current_time > expires_at:
609609
logger.info("Cleaning up expired file batch", id=info["id"])
610-
await self.metadata_store.delete(table=TABLE_VECTOR_STORE_FILE_BATCHES, where={"id": info["id"]})
610+
await self.metadata_store.sql_store.delete(
611+
table=TABLE_VECTOR_STORE_FILE_BATCHES,
612+
where={"id": info["id"]},
613+
)
611614
self.openai_file_batches.pop(info["id"], None)
612615
expired_count += 1
613616
if expired_count > 0:
@@ -723,6 +726,12 @@ async def initialize_openai_vector_stores(self) -> None:
723726
"Files API is not available. File attachment operations on vector stores will fail. "
724727
"Ensure a 'files' provider is configured if file operations are needed."
725728
)
729+
policy = getattr(self, "_policy", [])
730+
if policy and not self.metadata_store:
731+
raise ValueError(
732+
"Failed to initialize vector store provider: metadata_store is required when access control "
733+
"policies are configured. Configure storage.stores.vector_stores in your server config."
734+
)
726735
if self.metadata_store:
727736
await self._create_metadata_tables()
728737
if self.kvstore:
Lines changed: 123 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,123 @@
1+
# Copyright (c) The OGX Contributors.
2+
# All rights reserved.
3+
#
4+
# This source code is licensed under the terms described in the LICENSE file in
5+
# the root directory of this source tree.
6+
7+
from types import SimpleNamespace
8+
from typing import Any
9+
10+
from pydantic import BaseModel, Field
11+
12+
from ogx.core.datatypes import StackConfig
13+
from ogx.core.resolver import instantiate_provider
14+
from ogx.core.storage.datatypes import ServerStoresConfig, SqlStoreReference, StorageConfig
15+
from ogx_api import Api, RemoteProviderSpec
16+
17+
18+
class DummyRemoteConfig(BaseModel):
19+
metadata_store: SqlStoreReference | None = Field(default=None)
20+
21+
22+
def _make_run_config(vector_store_table: str = "vector_store_metadata") -> StackConfig:
23+
return StackConfig(
24+
distro_name="test",
25+
providers={},
26+
storage=StorageConfig(
27+
stores=ServerStoresConfig(
28+
vector_stores=SqlStoreReference(
29+
backend="sql_default",
30+
table_name=vector_store_table,
31+
)
32+
)
33+
),
34+
)
35+
36+
37+
class TestResolverMetadataStoreInjection:
38+
async def test_injects_metadata_store_for_remote_provider_when_missing(self, monkeypatch):
39+
captured: dict[str, Any] = {}
40+
41+
async def get_adapter_impl(config, deps, policy=None):
42+
captured["config"] = config
43+
return SimpleNamespace()
44+
45+
monkeypatch.setattr("ogx.core.resolver.instantiate_class_type", lambda _: DummyRemoteConfig)
46+
monkeypatch.setattr(
47+
"ogx.core.resolver.importlib.import_module",
48+
lambda _: SimpleNamespace(get_adapter_impl=get_adapter_impl),
49+
)
50+
monkeypatch.setattr("ogx.core.resolver.check_protocol_compliance", lambda *_args, **_kwargs: None)
51+
52+
provider = SimpleNamespace(
53+
provider_id="dummy-remote",
54+
provider_type="remote::dummy",
55+
config={},
56+
spec=RemoteProviderSpec(
57+
api=Api.vector_io,
58+
provider_type="remote::dummy",
59+
config_class="dummy.Config",
60+
module="dummy.remote.module",
61+
adapter_type="dummy-adapter",
62+
),
63+
)
64+
65+
await instantiate_provider(
66+
provider=provider,
67+
deps={},
68+
inner_impls={},
69+
dist_registry=SimpleNamespace(),
70+
run_config=_make_run_config(),
71+
policy=[],
72+
)
73+
74+
config = captured["config"]
75+
assert config.metadata_store is not None
76+
assert config.metadata_store.backend == "sql_default"
77+
assert config.metadata_store.table_name == "vector_store_metadata"
78+
79+
async def test_preserves_explicit_remote_metadata_store_config(self, monkeypatch):
80+
captured: dict[str, Any] = {}
81+
82+
async def get_adapter_impl(config, deps, policy=None):
83+
captured["config"] = config
84+
return SimpleNamespace()
85+
86+
monkeypatch.setattr("ogx.core.resolver.instantiate_class_type", lambda _: DummyRemoteConfig)
87+
monkeypatch.setattr(
88+
"ogx.core.resolver.importlib.import_module",
89+
lambda _: SimpleNamespace(get_adapter_impl=get_adapter_impl),
90+
)
91+
monkeypatch.setattr("ogx.core.resolver.check_protocol_compliance", lambda *_args, **_kwargs: None)
92+
93+
provider = SimpleNamespace(
94+
provider_id="dummy-remote",
95+
provider_type="remote::dummy",
96+
config={
97+
"metadata_store": {
98+
"backend": "sql_default",
99+
"table_name": "custom_vector_store_table",
100+
}
101+
},
102+
spec=RemoteProviderSpec(
103+
api=Api.vector_io,
104+
provider_type="remote::dummy",
105+
config_class="dummy.Config",
106+
module="dummy.remote.module",
107+
adapter_type="dummy-adapter",
108+
),
109+
)
110+
111+
await instantiate_provider(
112+
provider=provider,
113+
deps={},
114+
inner_impls={},
115+
dist_registry=SimpleNamespace(),
116+
run_config=_make_run_config(vector_store_table="default_vector_store_table"),
117+
policy=[],
118+
)
119+
120+
config = captured["config"]
121+
assert config.metadata_store is not None
122+
assert config.metadata_store.backend == "sql_default"
123+
assert config.metadata_store.table_name == "custom_vector_store_table"

tests/unit/providers/utils/memory/test_openai_vector_store_mixin.py

Lines changed: 89 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -327,3 +327,92 @@ async def test_migration_copies_batches(self):
327327
conflict_columns=["id"],
328328
update_columns=["store_id", "batch_data", "expires_at"],
329329
)
330+
331+
332+
class TestMetadataStoreEnforcement:
333+
"""Tests for mandatory metadata_store when access control policies are active."""
334+
335+
async def test_initialize_raises_when_policy_set_but_no_metadata_store(self):
336+
mixin = MockVectorStoreMixin(
337+
inference_api=AsyncMock(),
338+
files_api=AsyncMock(),
339+
kvstore=AsyncMock(),
340+
)
341+
mixin._policy = [MagicMock()]
342+
343+
with pytest.raises(ValueError, match="metadata_store is required"):
344+
await mixin.initialize_openai_vector_stores()
345+
346+
async def test_initialize_succeeds_when_no_policy(self):
347+
kvstore = AsyncMock()
348+
kvstore.values_in_range = AsyncMock(return_value=[])
349+
350+
mixin = MockVectorStoreMixin(
351+
inference_api=AsyncMock(),
352+
files_api=AsyncMock(),
353+
kvstore=kvstore,
354+
)
355+
mixin._policy = []
356+
357+
await mixin.initialize_openai_vector_stores()
358+
359+
async def test_initialize_succeeds_when_policy_and_metadata_store_set(self):
360+
metadata_store = MagicMock()
361+
metadata_store.create_table = AsyncMock()
362+
metadata_store.sql_store = AsyncMock()
363+
metadata_store.sql_store.fetch_all = AsyncMock(return_value=MagicMock(data=[]))
364+
metadata_store.fetch_all = AsyncMock(return_value=MagicMock(data=[]))
365+
366+
kvstore = AsyncMock()
367+
kvstore.get = AsyncMock(return_value="1")
368+
kvstore.values_in_range = AsyncMock(return_value=[])
369+
370+
mixin = MockVectorStoreMixin(
371+
inference_api=AsyncMock(),
372+
files_api=AsyncMock(),
373+
kvstore=kvstore,
374+
metadata_store=metadata_store,
375+
)
376+
mixin._policy = [MagicMock()]
377+
378+
await mixin.initialize_openai_vector_stores()
379+
380+
381+
class TestFileBatchCleanup:
382+
"""Tests for metadata-store file batch cleanup behavior."""
383+
384+
async def test_cleanup_uses_unfiltered_sql_store_access(self):
385+
sql_store = AsyncMock()
386+
sql_store.fetch_all = AsyncMock(
387+
return_value=MagicMock(
388+
data=[
389+
{
390+
"id": "batch_1",
391+
"batch_data": {"id": "batch_1", "expires_at": 1},
392+
}
393+
]
394+
)
395+
)
396+
sql_store.delete = AsyncMock()
397+
398+
metadata_store = MagicMock()
399+
metadata_store.sql_store = sql_store
400+
metadata_store.fetch_all = AsyncMock(side_effect=AssertionError("filtered fetch_all should not be used"))
401+
402+
mixin = MockVectorStoreMixin(
403+
inference_api=AsyncMock(),
404+
files_api=AsyncMock(),
405+
kvstore=AsyncMock(),
406+
metadata_store=metadata_store,
407+
)
408+
mixin.openai_file_batches = {"batch_1": {"id": "batch_1"}}
409+
410+
await mixin._cleanup_expired_file_batches()
411+
412+
metadata_store.fetch_all.assert_not_called()
413+
sql_store.fetch_all.assert_called_once_with(table="vector_store_file_batches")
414+
sql_store.delete.assert_called_once_with(
415+
table="vector_store_file_batches",
416+
where={"id": "batch_1"},
417+
)
418+
assert "batch_1" not in mixin.openai_file_batches

0 commit comments

Comments
 (0)