FIX: proper parameters in Astra DB Vectorize options (#3901)
* FIX: proper parameters in vectorize options * Update test_astra_component.py
This commit is contained in:
parent
b12fa9f874
commit
f403c17d10
2 changed files with 15 additions and 15 deletions
|
|
@ -322,20 +322,20 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
|
||||||
def build_vectorize_options(self, **kwargs):
|
def build_vectorize_options(self, **kwargs):
|
||||||
for attribute in [
|
for attribute in [
|
||||||
"provider",
|
"provider",
|
||||||
"z_00_api_key_name",
|
"z_00_model_name",
|
||||||
"z_01_model_name",
|
"z_01_model_parameters",
|
||||||
"z_02_authentication",
|
"z_02_api_key_name",
|
||||||
"z_03_provider_api_key",
|
"z_03_provider_api_key",
|
||||||
"z_04_model_parameters",
|
"z_04_authentication",
|
||||||
]:
|
]:
|
||||||
if not hasattr(self, attribute):
|
if not hasattr(self, attribute):
|
||||||
setattr(self, attribute, None)
|
setattr(self, attribute, None)
|
||||||
|
|
||||||
# Fetch values from kwargs if any self.* attributes are None
|
# Fetch values from kwargs if any self.* attributes are None
|
||||||
provider_value = self.VECTORIZE_PROVIDERS_MAPPING.get(self.provider, [None])[0] or kwargs.get("provider")
|
provider_value = self.VECTORIZE_PROVIDERS_MAPPING.get(self.provider, [None])[0] or kwargs.get("provider")
|
||||||
authentication = {**(self.z_02_authentication or kwargs.get("z_02_authentication", {}))}
|
authentication = {**(self.z_04_authentication or kwargs.get("z_04_authentication", {}))}
|
||||||
|
|
||||||
api_key_name = self.z_00_api_key_name or kwargs.get("z_00_api_key_name")
|
api_key_name = self.z_02_api_key_name or kwargs.get("z_02_api_key_name")
|
||||||
provider_key_name = self.z_03_provider_api_key or kwargs.get("z_03_provider_api_key")
|
provider_key_name = self.z_03_provider_api_key or kwargs.get("z_03_provider_api_key")
|
||||||
if provider_key_name:
|
if provider_key_name:
|
||||||
authentication["providerKey"] = provider_key_name
|
authentication["providerKey"] = provider_key_name
|
||||||
|
|
@ -346,9 +346,9 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
|
||||||
# must match astrapy.info.CollectionVectorServiceOptions
|
# must match astrapy.info.CollectionVectorServiceOptions
|
||||||
"collection_vector_service_options": {
|
"collection_vector_service_options": {
|
||||||
"provider": provider_value,
|
"provider": provider_value,
|
||||||
"modelName": self.z_01_model_name or kwargs.get("z_01_model_name"),
|
"modelName": self.z_00_model_name or kwargs.get("z_00_model_name"),
|
||||||
"authentication": authentication,
|
"authentication": authentication,
|
||||||
"parameters": self.z_04_model_parameters or kwargs.get("z_04_model_parameters", {}),
|
"parameters": self.z_01_model_parameters or kwargs.get("z_01_model_parameters", {}),
|
||||||
},
|
},
|
||||||
"collection_embedding_api_key": self.z_03_provider_api_key or kwargs.get("z_03_provider_api_key"),
|
"collection_embedding_api_key": self.z_03_provider_api_key or kwargs.get("z_03_provider_api_key"),
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -104,7 +104,7 @@ def test_astra_vectorize():
|
||||||
store = None
|
store = None
|
||||||
try:
|
try:
|
||||||
options = {"provider": "nvidia", "modelName": "NV-Embed-QA"}
|
options = {"provider": "nvidia", "modelName": "NV-Embed-QA"}
|
||||||
options_comp = {"provider": "nvidia", "z_01_model_name": "NV-Embed-QA"}
|
options_comp = {"provider": "nvidia", "z_00_model_name": "NV-Embed-QA"}
|
||||||
|
|
||||||
store = AstraDBVectorStore(
|
store = AstraDBVectorStore(
|
||||||
collection_name=VECTORIZE_COLLECTION,
|
collection_name=VECTORIZE_COLLECTION,
|
||||||
|
|
@ -156,10 +156,10 @@ def test_astra_vectorize_with_provider_api_key():
|
||||||
|
|
||||||
options_comp = {
|
options_comp = {
|
||||||
"provider": "openai",
|
"provider": "openai",
|
||||||
"z_01_model_name": "text-embedding-3-small",
|
"z_00_model_name": "text-embedding-3-small",
|
||||||
"z_04_model_parameters": {},
|
"z_01_model_parameters": {},
|
||||||
"z_02_authentication": {},
|
|
||||||
"z_03_provider_api_key": "openai",
|
"z_03_provider_api_key": "openai",
|
||||||
|
"z_04_authentication": {},
|
||||||
}
|
}
|
||||||
|
|
||||||
store = AstraDBVectorStore(
|
store = AstraDBVectorStore(
|
||||||
|
|
@ -212,9 +212,9 @@ def test_astra_vectorize_passes_authentication():
|
||||||
}
|
}
|
||||||
options_comp = {
|
options_comp = {
|
||||||
"provider": "openai",
|
"provider": "openai",
|
||||||
"z_01_model_name": "text-embedding-3-small",
|
"z_00_model_name": "text-embedding-3-small",
|
||||||
"z_04_model_parameters": {},
|
"z_01_model_parameters": {},
|
||||||
"z_02_authentication": {"providerKey": "openai"},
|
"z_04_authentication": {"providerKey": "openai"},
|
||||||
}
|
}
|
||||||
|
|
||||||
store = AstraDBVectorStore(
|
store = AstraDBVectorStore(
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue