fix(ts-oss, py): release Oracle client on init failure, validate insert batches (#6839)
This commit is contained in:
@@ -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)`,
|
||||
|
||||
@@ -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();
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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"),
|
||||
[
|
||||
|
||||
Reference in New Issue
Block a user