fix(ts-oss, py): release Oracle client on init failure, validate insert batches (#6839)

This commit is contained in:
Kartik
2026-08-06 18:16:25 +05:30
committed by GitHub
parent beea626f0a
commit 4a0a9a92a6
4 changed files with 145 additions and 20 deletions
+17 -2
View File
@@ -401,7 +401,15 @@ export class OracleAIVectorSearch implements VectorStore {
async initialize(): Promise<void> {
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<string, any>[],
): Promise<void> {
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)`,
+81 -1
View File
@@ -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();
});
});
+26 -17
View File
@@ -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):
+21
View File
@@ -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"),
[