From 4a0a9a92a641b5da75023eca596a758a4f7ae101 Mon Sep 17 00:00:00 2001 From: Kartik Date: Thu, 6 Aug 2026 18:16:25 +0530 Subject: [PATCH] fix(ts-oss, py): release Oracle client on init failure, validate insert batches (#6839) --- mem0-ts/src/oss/src/vector_stores/oracledb.ts | 19 ++++- mem0-ts/src/oss/tests/oracledb.unit.test.ts | 82 ++++++++++++++++++- mem0/vector_stores/oracledb.py | 43 ++++++---- tests/vector_stores/test_oracledb.py | 21 +++++ 4 files changed, 145 insertions(+), 20 deletions(-) diff --git a/mem0-ts/src/oss/src/vector_stores/oracledb.ts b/mem0-ts/src/oss/src/vector_stores/oracledb.ts index 85f7a2256..46cf016b8 100644 --- a/mem0-ts/src/oss/src/vector_stores/oracledb.ts +++ b/mem0-ts/src/oss/src/vector_stores/oracledb.ts @@ -401,7 +401,15 @@ export class OracleAIVectorSearch implements VectorStore { async initialize(): Promise { if (!this._initPromise) { - this._initPromise = this._doInitialize(); + this._initPromise = this._doInitialize().catch(async (error) => { + if (this.ownsClient && this.client) { + await Promise.resolve(this.client.close()).catch(() => {}); + this.client = undefined; + this.ownsClient = false; + } + this._initPromise = undefined; + throw error; + }); } return this._initPromise; } @@ -565,10 +573,17 @@ export class OracleAIVectorSearch implements VectorStore { ids: string[], payloads: Record[], ): Promise { - await this.initialize(); + if (ids.length !== vectors.length) { + throw new Error("ids and vectors must have the same length"); + } + if (payloads.length !== vectors.length) { + throw new Error("payloads and vectors must have the same length"); + } if (vectors.length === 0) return; + await this.initialize(); + await this.withConnection(async (connection) => { await connection.executeMany( `INSERT INTO ${this.collectionName} (id, vector, payload) VALUES (:id, :vector, :payload)`, diff --git a/mem0-ts/src/oss/tests/oracledb.unit.test.ts b/mem0-ts/src/oss/tests/oracledb.unit.test.ts index f1ea7b18f..3afdbac6e 100644 --- a/mem0-ts/src/oss/tests/oracledb.unit.test.ts +++ b/mem0-ts/src/oss/tests/oracledb.unit.test.ts @@ -4,9 +4,17 @@ const DB_TYPE_VECTOR = { name: "DB_TYPE_VECTOR" }; const DB_TYPE_JSON = { name: "DB_TYPE_JSON" }; const DB_TYPE_VARCHAR = { name: "DB_TYPE_VARCHAR" }; +const mockCreatePool = jest.fn(); + jest.mock( "oracledb", - () => ({ thin: true, DB_TYPE_VECTOR, DB_TYPE_JSON, DB_TYPE_VARCHAR }), + () => ({ + thin: true, + DB_TYPE_VECTOR, + DB_TYPE_JSON, + DB_TYPE_VARCHAR, + createPool: mockCreatePool, + }), { virtual: true }, ); @@ -297,6 +305,17 @@ describe("OracleAIVectorSearch SQL", () => { ); }); + it("rejects insert batches whose IDs or payloads do not match vectors", async () => { + const store = makeStore([]); + + await expect(store.insert([[1, 2, 3]], [], [{}])).rejects.toThrow( + "ids and vectors must have the same length", + ); + await expect(store.insert([[1, 2, 3]], ["id-1"], [])).rejects.toThrow( + "payloads and vectors must have the same length", + ); + }); + it("converts cosine distance to a similarity score", async () => { const calls: Call[] = []; const store = makeStore(calls, [[["id-1", { data: "hello" }, 0.25]]]); @@ -393,3 +412,64 @@ describe("OracleAIVectorSearch SQL", () => { expect(count).toBe(0); }); }); + +describe("OracleAIVectorSearch initialization failure", () => { + beforeEach(() => mockCreatePool.mockReset()); + + function fakePool(serverVersion: string) { + const connection = { + ...fakeConnection([], []), + oracleServerVersionString: serverVersion, + }; + return { + close: jest.fn(async () => {}), + getConnection: jest.fn(async () => connection), + }; + } + + function makePoolStore() { + return new OracleAIVectorSearch({ + connectionParams: { user: "u", password: "p", connectString: "d" }, + collectionName: "mem0", + embeddingModelDims: 3, + } as any); + } + + it("closes the pool it owns when initialization fails", async () => { + const pool = fakePool("23.3.0.24.05"); + mockCreatePool.mockResolvedValue(pool); + + await expect(makePoolStore().initialize()).rejects.toThrow( + "Oracle DB version 23.3.0.24.05 not supported", + ); + expect(pool.close).toHaveBeenCalledTimes(1); + }); + + it("does not cache the rejection, so a later attempt can succeed", async () => { + const failing = fakePool("23.3.0.24.05"); + const working = fakePool("23.4.0.24.05"); + mockCreatePool + .mockResolvedValueOnce(failing) + .mockResolvedValueOnce(working); + + const store = makePoolStore(); + await expect(store.initialize()).rejects.toThrow("not supported"); + await expect(store.initialize()).resolves.toBeUndefined(); + expect(mockCreatePool).toHaveBeenCalledTimes(2); + }); + + it("leaves a caller-supplied client open when initialization fails", async () => { + const client = fakeConnection([], []); + client.oracleServerVersionString = "23.3.0.24.05"; + const close = jest.spyOn(client, "close"); + + const store = new OracleAIVectorSearch({ + client: client as any, + collectionName: "mem0", + embeddingModelDims: 3, + } as any); + + await expect(store.initialize()).rejects.toThrow("not supported"); + expect(close).not.toHaveBeenCalled(); + }); +}); diff --git a/mem0/vector_stores/oracledb.py b/mem0/vector_stores/oracledb.py index 815baaed4..bf568bec7 100644 --- a/mem0/vector_stores/oracledb.py +++ b/mem0/vector_stores/oracledb.py @@ -245,25 +245,34 @@ class OracleAIVectorSearch(VectorStoreBase): self.client = oracledb.connect(**self.config.connection_params) self._owns_client = True - if not (hasattr(self.client, "thin") and self.client.thin): - if oracledb.clientversion()[:2] < (23, 4): - raise RuntimeError( - f"Oracle DB client driver version {'.'.join(map(str, oracledb.clientversion()))} " - "not supported, must be >=23.4 for vector support" + try: + if not (hasattr(self.client, "thin") and self.client.thin): + if oracledb.clientversion()[:2] < (23, 4): + raise RuntimeError( + f"Oracle DB client driver version {'.'.join(map(str, oracledb.clientversion()))} " + "not supported, must be >=23.4 for vector support" + ) + + if isinstance(self.client, oracledb.Connection): + db_version = tuple([int(v) for v in self.client.version.split(".")]) + else: + with self.client.acquire() as conn: + db_version = tuple([int(v) for v in conn.version.split(".")]) + + if db_version < (23, 4): + raise ValueError( + f"Oracle DB version {'.'.join(map(str, db_version))} not supported, " + "must be >=23.4 for vector support" ) - if isinstance(self.client, oracledb.Connection): - db_version = tuple([int(v) for v in self.client.version.split(".")]) - else: - with self.client.acquire() as conn: - db_version = tuple([int(v) for v in conn.version.split(".")]) - - if db_version < (23, 4): - raise ValueError( - f"Oracle DB version {'.'.join(map(str, db_version))} not supported, must be >=23.4 for vector support" - ) - - self.create_col() + self.create_col() + except Exception: + if self._owns_client: + try: + self.client.close() + except Exception: + pass + raise @contextmanager def _get_cursor(self, commit: bool = False): diff --git a/tests/vector_stores/test_oracledb.py b/tests/vector_stores/test_oracledb.py index 25ff9e25a..c2b4d87ed 100644 --- a/tests/vector_stores/test_oracledb.py +++ b/tests/vector_stores/test_oracledb.py @@ -331,6 +331,27 @@ def test_config_rejects_none_for_non_optional_fields(field, value): OracleAIVectorSearchConfig(client=object(), **{field: value}) +def test_init_closes_owned_client_when_post_connect_setup_fails(monkeypatch): + fake_connection = MagicMock(spec=oracledb.Connection) + fake_connection.thin = True + fake_connection.version = "23.4.0.0" + + monkeypatch.setattr(oracledb, "connect", MagicMock(return_value=fake_connection)) + monkeypatch.setattr(OracleAIVectorSearch, "create_col", MagicMock(side_effect=RuntimeError("boom"))) + monkeypatch.setattr(OracleAIVectorSearch, "__del__", lambda self: None) + + with pytest.raises(RuntimeError, match="boom"): + OracleAIVectorSearch( + collection_name=_unique_collection_name(), + embedding_model_dims=DIM, + connection_params={"user": "u", "password": "p", "dsn": "d"}, + use_connection_pool=False, + do_create_index=False, + ) + + fake_connection.close.assert_called_once() + + @pytest.mark.parametrize( ("metric", "distance", "expected_score"), [