diff --git a/tests/test_vikingdb_knowledge_backend.py b/tests/test_vikingdb_knowledge_backend.py index e8b5e7fcc..684a6d200 100644 --- a/tests/test_vikingdb_knowledge_backend.py +++ b/tests/test_vikingdb_knowledge_backend.py @@ -112,6 +112,8 @@ def test_byteplus_viking_knowledgebase_uses_hong_kong_fallback( assert ( backend.base_url == "https://api-knowledgebase.mlp.cn-hongkong.bytepluses.com" ) + assert backend.tos_config.region == "cn-hongkong" + assert backend.tos_config.endpoint == "tos-cn-hongkong.bytepluses.com" def test_byteplus_viking_knowledgebase_keeps_hong_kong_region( @@ -139,3 +141,161 @@ def test_byteplus_viking_knowledgebase_keeps_hong_kong_region( assert ( backend.base_url == "https://api-knowledgebase.mlp.cn-hongkong.bytepluses.com" ) + + +def test_byteplus_viking_knowledgebase_keeps_explicit_tos_config( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.configs.database_configs import NormalTOSConfig + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + monkeypatch.setenv("CLOUD_PROVIDER", "byteplus") + monkeypatch.setenv("AGENTKIT_CLOUD_PROVIDER", "byteplus") + monkeypatch.setenv("BYTEPLUS_ACCESS_KEY", "bp-ak") + monkeypatch.setenv("BYTEPLUS_SECRET_KEY", "bp-sk") + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "collection_status", + lambda self: {"existed": True}, + ) + + tos_config = NormalTOSConfig( + bucket="custom-bucket", + endpoint="tos-ap-southeast-1.bytepluses.com", + region="ap-southeast-1", + ) + backend = VikingDBKnowledgeBackend( + index="vikingkl_we4191n", + tos_config=tos_config, + ) + + assert backend.region == "cn-hongkong" + assert backend.tos_config.region == "ap-southeast-1" + assert backend.tos_config.endpoint == "tos-ap-southeast-1.bytepluses.com" + + +def test_byteplus_viking_knowledgebase_keeps_explicit_tosconfig( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.configs.database_configs import TOSConfig + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + monkeypatch.setenv("CLOUD_PROVIDER", "byteplus") + monkeypatch.setenv("AGENTKIT_CLOUD_PROVIDER", "byteplus") + monkeypatch.setenv("BYTEPLUS_ACCESS_KEY", "bp-ak") + monkeypatch.setenv("BYTEPLUS_SECRET_KEY", "bp-sk") + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "collection_status", + lambda self: {"existed": True}, + ) + + tos_config = TOSConfig( + endpoint="tos-ap-southeast-1.bytepluses.com", + region="ap-southeast-1", + ) + backend = VikingDBKnowledgeBackend( + index="vikingkl_we4191n", + tos_config=tos_config, + ) + + assert backend.region == "cn-hongkong" + assert backend.tos_config.region == "ap-southeast-1" + assert backend.tos_config.endpoint == "tos-ap-southeast-1.bytepluses.com" + + +def test_byteplus_viking_knowledgebase_keeps_explicit_tos_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + monkeypatch.setenv("CLOUD_PROVIDER", "byteplus") + monkeypatch.setenv("AGENTKIT_CLOUD_PROVIDER", "byteplus") + monkeypatch.setenv("DATABASE_TOS_REGION", "ap-southeast-1") + monkeypatch.setenv("DATABASE_TOS_ENDPOINT", "tos-ap-southeast-1.bytepluses.com") + monkeypatch.setenv("BYTEPLUS_ACCESS_KEY", "bp-ak") + monkeypatch.setenv("BYTEPLUS_SECRET_KEY", "bp-sk") + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "collection_status", + lambda self: {"existed": True}, + ) + + backend = VikingDBKnowledgeBackend(index="vikingkl_we4191n") + + assert backend.region == "cn-hongkong" + assert backend.tos_config.region == "ap-southeast-1" + assert backend.tos_config.endpoint == "tos-ap-southeast-1.bytepluses.com" + + +def test_byteplus_viking_knowledgebase_get_tos_client_uses_aligned_region( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import veadk.knowledgebase.backends.vikingdb_knowledge_backend as module + + monkeypatch.setenv("CLOUD_PROVIDER", "byteplus") + monkeypatch.setenv("AGENTKIT_CLOUD_PROVIDER", "byteplus") + monkeypatch.setenv("BYTEPLUS_ACCESS_KEY", "bp-ak") + monkeypatch.setenv("BYTEPLUS_SECRET_KEY", "bp-sk") + monkeypatch.setattr( + module.VikingDBKnowledgeBackend, + "collection_status", + lambda self: {"existed": True}, + ) + captured: dict[str, Any] = {} + + class _FakeVeTOS: + def __init__(self, **kwargs: Any) -> None: + captured.update(kwargs) + + monkeypatch.setattr(module, "VeTOS", _FakeVeTOS) + + backend = module.VikingDBKnowledgeBackend(index="vikingkl_we4191n") + backend._get_tos_client("kb-bucket") + + assert captured["region"] == "cn-hongkong" + assert captured["bucket_name"] == "kb-bucket" + + +def test_byteplus_viking_knowledgebase_bucket_creation_uses_aligned_region( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import veadk.configs.database_configs as config_module + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + monkeypatch.setenv("CLOUD_PROVIDER", "byteplus") + monkeypatch.setenv("AGENTKIT_CLOUD_PROVIDER", "byteplus") + monkeypatch.setenv("DATABASE_TOS_BUCKET", "kb-bucket") + monkeypatch.setenv("BYTEPLUS_ACCESS_KEY", "bp-ak") + monkeypatch.setenv("BYTEPLUS_SECRET_KEY", "bp-sk") + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "collection_status", + lambda self: {"existed": True}, + ) + captured: dict[str, Any] = {} + + class _FakeVeTOS: + def __init__(self, **kwargs: Any) -> None: + captured.update(kwargs) + + def create_bucket(self) -> bool: + captured["create_bucket_called"] = True + return True + + monkeypatch.setattr(config_module, "VeTOS", _FakeVeTOS) + + backend = VikingDBKnowledgeBackend(index="vikingkl_we4191n") + + assert backend.tos_config.bucket == "kb-bucket" + assert captured["region"] == "cn-hongkong" + assert captured["bucket_name"] == "kb-bucket" + assert captured["create_bucket_called"] is True diff --git a/veadk/configs/database_configs.py b/veadk/configs/database_configs.py index cb90a4bb6..c343b8899 100644 --- a/veadk/configs/database_configs.py +++ b/veadk/configs/database_configs.py @@ -164,16 +164,19 @@ class TOSConfig(BaseSettings): def model_post_init(self, __context, /) -> None: cloud_provider = os.getenv("CLOUD_PROVIDER", "volces").lower() + configured_fields = set(self.model_fields_set) if cloud_provider == "byteplus": - self.endpoint = "tos-ap-southeast-1.bytepluses.com" - self.region = "ap-southeast-1" + if "endpoint" not in configured_fields: + self.endpoint = "tos-ap-southeast-1.bytepluses.com" + if "region" not in configured_fields: + self.region = "ap-southeast-1" @cached_property def bucket(self) -> str: _bucket = os.getenv("DATABASE_TOS_BUCKET") or DEFAULT_TOS_BUCKET_NAME - VeTOS(bucket_name=_bucket).create_bucket() + VeTOS(region=self.region, bucket_name=_bucket).create_bucket() return _bucket diff --git a/veadk/knowledgebase/backends/vikingdb_knowledge_backend.py b/veadk/knowledgebase/backends/vikingdb_knowledge_backend.py index f0523493e..8b9b51eb8 100644 --- a/veadk/knowledgebase/backends/vikingdb_knowledge_backend.py +++ b/veadk/knowledgebase/backends/vikingdb_knowledge_backend.py @@ -22,10 +22,10 @@ import requests from pydantic import Field from typing_extensions import override -from volcengine.viking_knowledgebase import VikingKnowledgeBaseService from volcengine.auth.SignerV4 import SignerV4 from volcengine.base.Request import Request from volcengine.Credentials import Credentials +from volcengine.viking_knowledgebase import VikingKnowledgeBaseService import veadk.config # noqa E401 from veadk.auth.veauth.utils import ( @@ -33,12 +33,11 @@ get_credential_from_vefaas_iam, ) from veadk.configs.database_configs import NormalTOSConfig, TOSConfig +from veadk.integrations.ve_tos.ve_tos import VeTOS from veadk.knowledgebase.backends.base_backend import BaseKnowledgebaseBackend from veadk.knowledgebase.entry import KnowledgebaseEntry from veadk.utils.logger import get_logger from veadk.utils.misc import formatted_timestamp, getenv -from veadk.integrations.ve_tos.ve_tos import VeTOS - logger = get_logger(__name__) @@ -231,6 +230,13 @@ def model_post_init(self, __context: Any) -> None: self.host = ( self.host or f"api-knowledgebase.mlp.{self.region}.bytepluses.com" ) + if ( + "tos_config" not in self.model_fields_set + and not os.getenv("DATABASE_TOS_REGION") + and not os.getenv("DATABASE_TOS_ENDPOINT") + ): + self.tos_config.region = self.region + self.tos_config.endpoint = f"tos-{self.region}.bytepluses.com" elif not self.region: self.region = os.getenv("DATABASE_VIKING_REGION", "cn-beijing") self.base_url = f"https://api-knowledgebase.mlp.{self.region}.volces.com"