From ac71516bf33c593cb5e669285d1139bfa6745b34 Mon Sep 17 00:00:00 2001 From: shanchunhua Date: Tue, 15 Sep 2026 15:23:03 +0800 Subject: [PATCH] feat: support Viking API key auth --- config.yaml.full | 6 + .../framework/knowledgebase/overview.en.mdx | 7 +- .../docs/framework/knowledgebase/overview.mdx | 7 +- .../memory/long-term/vikingdb.en.mdx | 3 +- .../framework/memory/long-term/vikingdb.mdx | 3 +- .../environment-variables.en.mdx | 4 + .../configuration/environment-variables.mdx | 4 + pyproject.toml | 2 +- tests/test_vikingdb_knowledge_backend.py | 296 ++++++++++++++++++ tests/test_vikingdb_memory_backend.py | 261 +++++++++++++++ uv.lock | 2 +- veadk/configs/database_configs.py | 3 + .../ve_viking_db_memory.py | 7 + .../backends/vikingdb_knowledge_backend.py | 193 ++++++++---- .../vikingdb_memory_backend.py | 156 +++++---- 15 files changed, 830 insertions(+), 124 deletions(-) diff --git a/config.yaml.full b/config.yaml.full index 05df99dc7..196cc143e 100644 --- a/config.yaml.full +++ b/config.yaml.full @@ -168,8 +168,14 @@ database: # output_fields: text,metadata # [optional] for knowledgebase (https://console.volcengine.com/vikingdb) viking: + # api_key: ${DATABASE_VIKING_API_KEY} # optional; for searching existing VikingDB knowledgebase collections project: default # user project in Volcengine Viking DB region: cn-beijing + # [optional] for LongTermMemory(backend="viking") + vikingmem: + # api_key: ${DATABASE_VIKINGMEM_API_KEY} # optional; for memory operations on existing VikingDB memory collections + project: default + memory_type: sys_event_v1,sys_profile_v1 # [optional] for knowledgebase with viking database tos: endpoint: tos-cn-beijing.volces.com # default Volcengine TOS endpoint diff --git a/docs/content/docs/framework/knowledgebase/overview.en.mdx b/docs/content/docs/framework/knowledgebase/overview.en.mdx index 110253bd8..181ac8c06 100644 --- a/docs/content/docs/framework/knowledgebase/overview.en.mdx +++ b/docs/content/docs/framework/knowledgebase/overview.en.mdx @@ -136,12 +136,14 @@ model: api_key: ${MODEL_AGENT_API_KEY} volcengine: - # Viking DB and embedding require Volcengine ak/sk + # Viking DB and embedding can use Volcengine ak/sk access_key: ${VOLCENGINE_ACCESS_KEY} secret_key: ${VOLCENGINE_SECRET_KEY} database: viking: + # Optional; used to search existing VikingDB knowledgebase collections + api_key: ${DATABASE_VIKING_API_KEY} project: default # project in Volcengine Viking DB region: cn-beijing tos: @@ -167,12 +169,15 @@ Provide those secrets in environment variables: ```bash export MODEL_AGENT_API_KEY="" +export DATABASE_VIKING_API_KEY="" export VOLCENGINE_ACCESS_KEY="" export VOLCENGINE_SECRET_KEY="" export DATABASE_OPENSEARCH_PASSWORD="" export DATABASE_OPENVIKING_API_KEY="" ``` +`DATABASE_VIKING_API_KEY` is preferred for searching existing VikingDB knowledgebase collections. Creating, deleting, or listing collections, managing documents/chunks, and operations such as `add_from_text` / `add_from_files` that upload through TOS still require usable `VOLCENGINE_ACCESS_KEY` / `VOLCENGINE_SECRET_KEY` or VeFaaS IAM credentials. + ### Demo scenario The example below uses short bios of two psychologists as knowledge, plus a date-difference tool, to demonstrate "retrieve from the knowledge base + call a tool." diff --git a/docs/content/docs/framework/knowledgebase/overview.mdx b/docs/content/docs/framework/knowledgebase/overview.mdx index 18c6fc5b2..465da9b5a 100644 --- a/docs/content/docs/framework/knowledgebase/overview.mdx +++ b/docs/content/docs/framework/knowledgebase/overview.mdx @@ -136,12 +136,14 @@ model: api_key: ${MODEL_AGENT_API_KEY} volcengine: - # Viking DB 与 embedding 需要火山引擎 ak/sk + # Viking DB 与 embedding 可使用火山引擎 ak/sk access_key: ${VOLCENGINE_ACCESS_KEY} secret_key: ${VOLCENGINE_SECRET_KEY} database: viking: + # 可选;用于搜索已有 VikingDB 知识库 + api_key: ${DATABASE_VIKING_API_KEY} project: default # Volcengine Viking DB 中的项目 region: cn-beijing tos: @@ -167,12 +169,15 @@ database: ```bash export MODEL_AGENT_API_KEY="<你的方舟 API Key>" +export DATABASE_VIKING_API_KEY="<你的 VikingDB 知识库 API Key>" export VOLCENGINE_ACCESS_KEY="<你的 AK>" export VOLCENGINE_SECRET_KEY="<你的 SK>" export DATABASE_OPENSEARCH_PASSWORD="<你的 OpenSearch 密码>" export DATABASE_OPENVIKING_API_KEY="<你的 OpenViking API Key>" ``` +`DATABASE_VIKING_API_KEY` 会优先用于搜索已有 VikingDB 知识库。创建、删除、列举 collection,管理文档/切片,以及 `add_from_text` / `add_from_files` 这类需要上传 TOS 的能力,仍需要可用的 `VOLCENGINE_ACCESS_KEY` / `VOLCENGINE_SECRET_KEY` 或 VeFaaS IAM 凭证。 + ### 演示场景 下面用一段心理学家简介作为知识,并配一个计算日期差的工具,演示「检索知识库 + 调用工具」。 diff --git a/docs/content/docs/framework/memory/long-term/vikingdb.en.mdx b/docs/content/docs/framework/memory/long-term/vikingdb.en.mdx index 3f16631fa..474179475 100644 --- a/docs/content/docs/framework/memory/long-term/vikingdb.en.mdx +++ b/docs/content/docs/framework/memory/long-term/vikingdb.en.mdx @@ -11,10 +11,11 @@ The `viking` backend (`VikingDBLTMBackend`) uses Volcengine [VikingDB memory](ht ## Configuration -Authenticated via Volcengine AK/SK. VeADK reads AK/SK from environment variables first, falling back to VeFaaS IAM credentials (for cloud deployments). +Authenticate with a VikingDB memory API key or Volcengine AK/SK. VeADK prefers `api_key` / `DATABASE_VIKINGMEM_API_KEY` for memory operations on an existing collection; when it is unset, it uses AK/SK and then falls back to VeFaaS IAM credentials for cloud deployments. Collection creation, listing, and deletion still require AK/SK or IAM credentials. | Field | Env var | Default | Description | | :--- | :--- | :--- | :--- | +| `api_key` | `DATABASE_VIKINGMEM_API_KEY` | — | VikingDB memory API key for memory operations on an existing collection. | | `volcengine_access_key` | `VOLCENGINE_ACCESS_KEY` | — | Volcengine Access Key. | | `volcengine_secret_key` | `VOLCENGINE_SECRET_KEY` | — | Volcengine Secret Key. | | `volcengine_project` | `DATABASE_VIKINGMEM_PROJECT` | `default` | VikingDB memory project. | diff --git a/docs/content/docs/framework/memory/long-term/vikingdb.mdx b/docs/content/docs/framework/memory/long-term/vikingdb.mdx index 464934a54..aee50d070 100644 --- a/docs/content/docs/framework/memory/long-term/vikingdb.mdx +++ b/docs/content/docs/framework/memory/long-term/vikingdb.mdx @@ -11,10 +11,11 @@ title: "VikingDB(viking)" ## 配置 -通过火山引擎 AK/SK 鉴权。VeADK 会优先从环境变量读取 AK/SK,未设置时尝试从 VeFaaS IAM 获取凭证(云上部署场景)。 +可以通过 VikingDB 记忆库 API Key 或火山引擎 AK/SK 鉴权。VeADK 会优先使用 `api_key` / `DATABASE_VIKINGMEM_API_KEY` 访问已有 collection 中的记忆数据;未设置时使用 AK/SK,仍未设置时尝试从 VeFaaS IAM 获取凭证(云上部署场景)。创建、列举、删除 collection 等管理操作仍需要 AK/SK 或 IAM 凭证。 | 配置项 | 环境变量 | 默认值 | 说明 | | :--- | :--- | :--- | :--- | +| `api_key` | `DATABASE_VIKINGMEM_API_KEY` | — | VikingDB 记忆库 API Key,用于访问已有 collection 中的记忆数据。 | | `volcengine_access_key` | `VOLCENGINE_ACCESS_KEY` | — | 火山引擎 Access Key。 | | `volcengine_secret_key` | `VOLCENGINE_SECRET_KEY` | — | 火山引擎 Secret Key。 | | `volcengine_project` | `DATABASE_VIKINGMEM_PROJECT` | `default` | VikingDB 记忆库项目。 | diff --git a/docs/content/docs/references/configuration/environment-variables.en.mdx b/docs/content/docs/references/configuration/environment-variables.en.mdx index f9bdc9977..0207053df 100644 --- a/docs/content/docs/references/configuration/environment-variables.en.mdx +++ b/docs/content/docs/references/configuration/environment-variables.en.mdx @@ -108,8 +108,12 @@ Prefix `DATABASE_`, grouped by storage type. Memory and the knowledge base read | | `DATABASE_MILVUS_OVERWRITE` | Overwrite the collection, default `false` | | | `DATABASE_MILVUS_TIMEOUT` | Connection/request timeout (optional) | | | `DATABASE_MILVUS_OUTPUT_FIELDS` | Comma-separated retrieval output fields (optional) | +| VikingDB knowledgebase | `DATABASE_VIKING_API_KEY` | API key for searching existing knowledgebases; management still requires AK/SK/IAM | | VikingDB | `DATABASE_VIKING_PROJECT` | Project name | | | `DATABASE_VIKING_REGION` | Region | +| VikingDB memory | `DATABASE_VIKINGMEM_API_KEY` | API key for existing memory collections; management still requires AK/SK/IAM | +| | `DATABASE_VIKINGMEM_PROJECT` | Project name, default `default` | +| | `DATABASE_VIKINGMEM_MEMORY_TYPE` | Comma-separated memory types | | TOS | `DATABASE_TOS_ENDPOINT` | Endpoint | | | `DATABASE_TOS_REGION` | Region | | | `DATABASE_TOS_BUCKET` | Bucket name | diff --git a/docs/content/docs/references/configuration/environment-variables.mdx b/docs/content/docs/references/configuration/environment-variables.mdx index 4dd64cfd3..d1209350b 100644 --- a/docs/content/docs/references/configuration/environment-variables.mdx +++ b/docs/content/docs/references/configuration/environment-variables.mdx @@ -108,8 +108,12 @@ volcengine: | | `DATABASE_MILVUS_OVERWRITE` | 是否覆盖 collection,默认 `false` | | | `DATABASE_MILVUS_TIMEOUT` | 连接/请求超时时间(可选) | | | `DATABASE_MILVUS_OUTPUT_FIELDS` | 检索输出字段,逗号分隔(可选) | +| VikingDB 知识库 | `DATABASE_VIKING_API_KEY` | API Key,用于搜索已有知识库;管理操作仍需 AK/SK/IAM | | VikingDB | `DATABASE_VIKING_PROJECT` | 项目名称 | | | `DATABASE_VIKING_REGION` | 区域 | +| VikingDB 记忆库 | `DATABASE_VIKINGMEM_API_KEY` | API Key,用于访问已有记忆 collection;管理操作仍需 AK/SK/IAM | +| | `DATABASE_VIKINGMEM_PROJECT` | 项目名称,默认 `default` | +| | `DATABASE_VIKINGMEM_MEMORY_TYPE` | 记忆类型,逗号分隔 | | TOS | `DATABASE_TOS_ENDPOINT` | 访问端点 | | | `DATABASE_TOS_REGION` | 区域 | | | `DATABASE_TOS_BUCKET` | 桶名称 | diff --git a/pyproject.toml b/pyproject.toml index 013069cda..a067e6551 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,7 +52,7 @@ dependencies = [ "pypdfium2>=4.30.0", # Frontend PDF attachments are rendered as page images "pillow>=10.0.0", # Encode rendered PDF pages as PNG "qrcode>=8,<9", # Render Feishu app-registration links for Studio setup - "vikingdb-python-sdk>=0.1.3", # For Viking DB + "vikingdb-python-sdk>=0.1.32", # For Viking DB "agentkit-sdk-python>=0.8.0", "websockets>=15,<16", # Stream temporary AgentKit Sandbox conversations "openviking-sdk>=0.1.3", diff --git a/tests/test_vikingdb_knowledge_backend.py b/tests/test_vikingdb_knowledge_backend.py index c814725dc..64614a1cb 100644 --- a/tests/test_vikingdb_knowledge_backend.py +++ b/tests/test_vikingdb_knowledge_backend.py @@ -14,6 +14,7 @@ from __future__ import annotations +from types import SimpleNamespace from typing import Any import pytest @@ -87,6 +88,301 @@ def test_viking_knowledgebase_reads_byteplus_credentials( assert backend.session_token == "bp-token" +def test_viking_knowledgebase_agentkit_provider_selects_byteplus_endpoint( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + monkeypatch.setenv("AGENTKIT_CLOUD_PROVIDER", "byteplus") + monkeypatch.delenv("CLOUD_PROVIDER", raising=False) + 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.cloud_provider == "byteplus" + assert backend.region == "cn-hongkong" + assert backend.host == "api-knowledgebase.mlp.cn-hongkong.bytepluses.com" + + +def test_viking_knowledgebase_reads_api_key_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + monkeypatch.setenv("DATABASE_VIKING_API_KEY", "kb-api-key") + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "model_post_init", + lambda self, __context: None, + ) + + backend = VikingDBKnowledgeBackend(index="vikingkl_we4191n") + + assert backend.api_key == "kb-api-key" + + +def test_viking_knowledgebase_ignores_empty_yaml_api_key( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + monkeypatch.setenv("DATABASE_VIKING_API_KEY", "None") + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "collection_status", + lambda self: {"existed": True}, + ) + + backend = VikingDBKnowledgeBackend(index="vikingkl_we4191n") + + assert backend.api_key is None + + +def test_viking_knowledgebase_empty_api_key_parameter_falls_back_to_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + monkeypatch.setenv("DATABASE_VIKING_API_KEY", "kb-api-key") + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "collection_status", + lambda self: {"existed": True}, + ) + + backend = VikingDBKnowledgeBackend(index="vikingkl_we4191n", api_key="") + + assert backend.api_key == "kb-api-key" + + +def test_viking_knowledgebase_explicit_region_sets_default_endpoint( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "collection_status", + lambda self: {"existed": True}, + ) + + backend = VikingDBKnowledgeBackend( + index="vikingkl_we4191n", + api_key="kb-api-key", + region="cn-shanghai", + ) + + assert backend.region == "cn-shanghai" + assert backend.base_url == "https://api-knowledgebase.mlp.cn-shanghai.volces.com" + assert backend.host == "api-knowledgebase.mlp.cn-shanghai.volces.com" + + +def test_viking_knowledgebase_do_request_uses_api_key_bearer( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.knowledgebase.backends import vikingdb_knowledge_backend as module + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + captured: dict[str, Any] = {} + + class _FakeResponse: + ok = True + + def json(self) -> dict[str, Any]: + return {"code": 0, "data": {"ok": True}} + + def fake_request(**kwargs: Any) -> _FakeResponse: + captured.update(kwargs) + return _FakeResponse() + + def unexpected_signed_request(*_: Any, **__: Any) -> None: + raise AssertionError("API key auth must not build an AK/SK signed request") + + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "model_post_init", + lambda self, __context: None, + ) + monkeypatch.setattr(module.requests, "request", fake_request) + monkeypatch.setattr( + module, + "build_vikingdb_knowledgebase_request", + unexpected_signed_request, + ) + + backend = VikingDBKnowledgeBackend( + index="vikingkl_we4191n", + api_key="kb-api-key", + base_url="https://viking.example", + ) + result = backend._do_request( + body={"name": "vikingkl_we4191n"}, + path="/api/knowledge/collection/info", + auth_mode="api_key", + ) + + assert result == {"code": 0, "data": {"ok": True}} + assert captured["url"] == "https://viking.example/api/knowledge/collection/info" + assert captured["headers"]["Authorization"] == "Bearer kb-api-key" + assert captured["data"] == '{"name": "vikingkl_we4191n"}' + + +def test_viking_knowledgebase_management_request_uses_ak_sk_with_api_key( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.knowledgebase.backends import vikingdb_knowledge_backend as module + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + captured: dict[str, Any] = {} + + class _FakeResponse: + ok = True + + def json(self) -> dict[str, Any]: + return {"code": 0, "data": {"ok": True}} + + def fake_signed_request(**kwargs: Any) -> SimpleNamespace: + captured["signed_request"] = kwargs + return SimpleNamespace( + headers={"Authorization": "Signed AK/SK"}, + body='{"signed": true}', + ) + + def fake_request(**kwargs: Any) -> _FakeResponse: + captured["request"] = kwargs + return _FakeResponse() + + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "model_post_init", + lambda self, __context: None, + ) + monkeypatch.setattr( + module, + "build_vikingdb_knowledgebase_request", + fake_signed_request, + ) + monkeypatch.setattr(module.requests, "request", fake_request) + + backend = VikingDBKnowledgeBackend( + index="vikingkl_we4191n", + api_key="kb-api-key", + volcengine_access_key="ak", + volcengine_secret_key="sk", + base_url="https://viking.example", + ) + result = backend._do_request( + body={"name": "vikingkl_we4191n"}, + path="/api/knowledge/collection/info", + ) + + assert result == {"code": 0, "data": {"ok": True}} + assert captured["signed_request"]["volcengine_access_key"] == "ak" + assert captured["signed_request"]["volcengine_secret_key"] == "sk" + assert captured["request"]["headers"]["Authorization"] == "Signed AK/SK" + assert captured["request"]["headers"]["Authorization"] != "Bearer kb-api-key" + + +def test_viking_knowledgebase_api_key_only_skips_collection_management_precheck( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + monkeypatch.setenv("DATABASE_VIKING_API_KEY", "kb-api-key") + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "collection_status", + lambda self: (_ for _ in ()).throw( + AssertionError("API-key-only init must not check collection management") + ), + ) + + backend = VikingDBKnowledgeBackend( + index="vikingkl_we4191n", + volcengine_access_key=None, + volcengine_secret_key=None, + ) + + assert backend.api_key == "kb-api-key" + + +def test_viking_knowledgebase_search_uses_api_key_request( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.knowledgebase.backends.vikingdb_knowledge_backend import ( + VikingDBKnowledgeBackend, + ) + + captured: dict[str, Any] = {} + + def fake_do_request(self, **kwargs: Any) -> dict[str, Any]: + del self + captured.update(kwargs) + return { + "code": 0, + "data": { + "rewrite_query": None, + "result_list": [ + { + "content": "answer", + "doc_info": { + "doc_meta": '[{"field_name": "source", "field_value": "manual"}]' + }, + } + ], + }, + } + + monkeypatch.setattr( + VikingDBKnowledgeBackend, + "model_post_init", + lambda self, __context: None, + ) + monkeypatch.setattr(VikingDBKnowledgeBackend, "_do_request", fake_do_request) + + backend = VikingDBKnowledgeBackend( + index="vikingkl_we4191n", + api_key="kb-api-key", + resource_id="kb-yef-example", + ) + entries = backend.search("hello", metadata={"source": "manual"}) + + assert captured["path"] == "/api/knowledge/collection/search_knowledge" + assert captured["auth_mode"] == "api_key" + assert captured["body"]["resource_id"] == "kb-yef-example" + assert captured["body"]["query_param"] == { + "doc_filter": { + "op": "and", + "conds": [{"op": "must", "field": "source", "conds": ["manual"]}], + } + } + assert len(entries) == 1 + assert entries[0].content == "answer" + assert entries[0].metadata == {"source": "manual"} + + def test_byteplus_viking_knowledgebase_uses_hong_kong_fallback( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/tests/test_vikingdb_memory_backend.py b/tests/test_vikingdb_memory_backend.py index f866f2dcd..9fea683e1 100644 --- a/tests/test_vikingdb_memory_backend.py +++ b/tests/test_vikingdb_memory_backend.py @@ -23,11 +23,16 @@ def _load_backend_module(monkeypatch: pytest.MonkeyPatch): + class FakeAPIKey: + def __init__(self, **kwargs: Any) -> None: + self.kwargs = kwargs + class FakeIAM: def __init__(self, **_: Any) -> None: pass class FakeVikingDBModule(types.ModuleType): + APIKey: type[FakeAPIKey] IAM: type[FakeIAM] class FakeVikingMem: @@ -38,6 +43,7 @@ class FakeVikingDBMemoryModule(types.ModuleType): VikingMem: type[FakeVikingMem] vikingdb_module = FakeVikingDBModule("vikingdb") + vikingdb_module.APIKey = FakeAPIKey vikingdb_module.IAM = FakeIAM vikingdb_memory_module = FakeVikingDBMemoryModule("vikingdb.memory") vikingdb_memory_module.VikingMem = FakeVikingMem @@ -69,6 +75,215 @@ def test_byteplus_viking_memory_uses_fixed_hong_kong_region( assert backend.region == "cn-hongkong" +def test_viking_memory_agentkit_provider_selects_byteplus_region( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = _load_backend_module(monkeypatch) + monkeypatch.setenv("AGENTKIT_CLOUD_PROVIDER", "byteplus") + monkeypatch.delenv("CLOUD_PROVIDER", raising=False) + monkeypatch.setenv("BYTEPLUS_ACCESS_KEY", "bp-ak") + monkeypatch.setenv("BYTEPLUS_SECRET_KEY", "bp-sk") + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_collection_exist", + lambda self: True, + ) + + backend = module.VikingDBLTMBackend(index="agent_memory") + + assert backend.cloud_provider == "byteplus" + assert backend.region == "cn-hongkong" + + +def test_viking_memory_reads_api_key_env(monkeypatch: pytest.MonkeyPatch) -> None: + module = _load_backend_module(monkeypatch) + monkeypatch.setenv("DATABASE_VIKINGMEM_API_KEY", "mem-api-key") + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_collection_exist", + lambda self: True, + ) + + backend = module.VikingDBLTMBackend(index="agent_memory") + + assert backend.api_key == "mem-api-key" + + +def test_viking_memory_sdk_client_uses_api_key_auth( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = _load_backend_module(monkeypatch) + captured: dict[str, Any] = {} + + class UnexpectedClient: + def __init__(self, **kwargs: Any) -> None: + raise AssertionError("API key SDK client must not use management client") + + class FakeAPIKey: + def __init__(self, **kwargs: Any) -> None: + captured["api_key_kwargs"] = kwargs + + class FakeVikingMem: + def __init__(self, **kwargs: Any) -> None: + captured["viking_mem_kwargs"] = kwargs + + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_collection_exist", + lambda self: True, + ) + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_get_ak_sk_sts", + lambda self: (_ for _ in ()).throw(AssertionError("AK/SK not expected")), + ) + monkeypatch.setattr(module, "VikingDBMemoryClient", UnexpectedClient) + monkeypatch.setattr(module, "APIKey", FakeAPIKey) + monkeypatch.setattr(module, "VikingMem", FakeVikingMem) + monkeypatch.setenv("DATABASE_VIKINGMEM_BASE_URL", "http://memory.example") + + backend = module.VikingDBLTMBackend( + index="agent_memory", + api_key="mem-api-key", + volcengine_access_key=None, + volcengine_secret_key=None, + ) + backend._get_sdk_client() + + assert captured["api_key_kwargs"] == {"api_key": "mem-api-key"} + assert isinstance(captured["viking_mem_kwargs"]["auth"], FakeAPIKey) + assert captured["viking_mem_kwargs"]["host"] == "memory.example" + assert captured["viking_mem_kwargs"]["scheme"] == "http" + + +def test_viking_memory_api_key_only_skips_collection_management_precheck( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = _load_backend_module(monkeypatch) + monkeypatch.setenv("DATABASE_VIKINGMEM_API_KEY", "mem-api-key") + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_collection_exist", + lambda self: (_ for _ in ()).throw( + AssertionError("API-key-only init must not check collection management") + ), + ) + + backend = module.VikingDBLTMBackend( + index="agent_memory", + volcengine_access_key=None, + volcengine_secret_key=None, + ) + + assert backend.api_key == "mem-api-key" + + +def test_viking_memory_get_user_profile_uses_sdk_search( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = _load_backend_module(monkeypatch) + captured: dict[str, Any] = {} + + class FakeCollection: + def search_memory(self, **kwargs: Any) -> dict[str, Any]: + captured["search_kwargs"] = kwargs + return { + "code": 0, + "data": { + "result_list": [ + {"memory_info": {"user_profile": "likes concise answers"}} + ] + }, + } + + class FakeSdkClient: + def get_collection(self, **kwargs: Any) -> FakeCollection: + captured["get_collection_kwargs"] = kwargs + return FakeCollection() + + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_collection_exist", + lambda self: True, + ) + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_get_sdk_client", + lambda self: FakeSdkClient(), + ) + + backend = module.VikingDBLTMBackend(index="agent_memory", api_key="mem-api-key") + + assert backend.get_user_profile("user-42") == "likes concise answers" + assert captured["get_collection_kwargs"] == { + "collection_name": "agent_memory", + "project_name": "default", + } + assert captured["search_kwargs"] == { + "filter": {"user_id": ["user-42"], "memory_category": 1}, + "limit": 5000, + } + + +def test_viking_memory_get_user_profile_raises_on_error_code( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = _load_backend_module(monkeypatch) + + class FakeCollection: + def search_memory(self, **kwargs: Any) -> dict[str, Any]: + return {"code": 1000001, "message": "unauthorized"} + + class FakeSdkClient: + def get_collection(self, **kwargs: Any) -> FakeCollection: + return FakeCollection() + + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_collection_exist", + lambda self: True, + ) + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_get_sdk_client", + lambda self: FakeSdkClient(), + ) + + backend = module.VikingDBLTMBackend(index="agent_memory", api_key="mem-api-key") + + with pytest.raises(ValueError, match="Get VikingDB user profile error"): + backend.get_user_profile("user-42") + + +def test_viking_memory_get_user_profile_returns_empty_when_not_found( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = _load_backend_module(monkeypatch) + + class FakeCollection: + def search_memory(self, **kwargs: Any) -> dict[str, Any]: + return {"code": 0, "data": {"result_list": []}} + + class FakeSdkClient: + def get_collection(self, **kwargs: Any) -> FakeCollection: + return FakeCollection() + + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_collection_exist", + lambda self: True, + ) + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_get_sdk_client", + lambda self: FakeSdkClient(), + ) + + backend = module.VikingDBLTMBackend(index="agent_memory", api_key="mem-api-key") + + assert backend.get_user_profile("user-42") == "" + + def test_byteplus_viking_memory_ignores_explicit_region( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -177,6 +392,52 @@ def fake_json(self, api: str, params: dict[str, Any], body: str) -> str: assert result["Result"]["Collections"][0]["CollectionName"] == "agent_memory" +def test_direct_viking_memory_client_rejects_api_key_for_management( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from veadk.integrations.ve_viking_db_memory.ve_viking_db_memory import ( + VikingDBMemoryClient, + ) + + if hasattr(VikingDBMemoryClient, "_instance"): + delattr(VikingDBMemoryClient, "_instance") + + with pytest.raises(ValueError, match="collection management requires AK/SK"): + VikingDBMemoryClient(region="cn-beijing", api_key="mem-api-key") + + +def test_viking_memory_ignores_empty_yaml_api_key( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = _load_backend_module(monkeypatch) + monkeypatch.setenv("DATABASE_VIKINGMEM_API_KEY", "None") + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_collection_exist", + lambda self: True, + ) + + backend = module.VikingDBLTMBackend(index="agent_memory") + + assert backend.api_key is None + + +def test_viking_memory_empty_api_key_parameter_falls_back_to_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = _load_backend_module(monkeypatch) + monkeypatch.setenv("DATABASE_VIKINGMEM_API_KEY", "mem-api-key") + monkeypatch.setattr( + module.VikingDBLTMBackend, + "_collection_exist", + lambda self: True, + ) + + backend = module.VikingDBLTMBackend(index="agent_memory", api_key="") + + assert backend.api_key == "mem-api-key" + + def test_byteplus_viking_memory_keeps_hong_kong_region( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/uv.lock b/uv.lock index 2b4a439a2..373c15611 100644 --- a/uv.lock +++ b/uv.lock @@ -5957,7 +5957,7 @@ requires-dist = [ { name = "trafilatura", specifier = ">=2.0,<2.1" }, { name = "trustedmcp", specifier = "==0.0.5" }, { name = "uvicorn", marker = "extra == 'codex'" }, - { name = "vikingdb-python-sdk", specifier = ">=0.1.3" }, + { name = "vikingdb-python-sdk", specifier = ">=0.1.32" }, { name = "volcengine", specifier = ">=1.0.193" }, { name = "volcengine", marker = "extra == 'database'", specifier = ">=1.0.193" }, { name = "volcengine-python-sdk", specifier = ">=5.0.36" }, diff --git a/veadk/configs/database_configs.py b/veadk/configs/database_configs.py index ba89cfb6f..7023a1222 100644 --- a/veadk/configs/database_configs.py +++ b/veadk/configs/database_configs.py @@ -162,6 +162,9 @@ class OpenVikingConfig(BaseSettings): class VikingKnowledgebaseConfig(BaseSettings): model_config = SettingsConfigDict(env_prefix="DATABASE_VIKING_") + api_key: str = "" + """VikingDB knowledgebase API key for searching existing collections.""" + project: str = "default" """User project in Volcengine console web.""" diff --git a/veadk/integrations/ve_viking_db_memory/ve_viking_db_memory.py b/veadk/integrations/ve_viking_db_memory/ve_viking_db_memory.py index c4cb536f9..5cff73b43 100644 --- a/veadk/integrations/ve_viking_db_memory/ve_viking_db_memory.py +++ b/veadk/integrations/ve_viking_db_memory/ve_viking_db_memory.py @@ -60,6 +60,7 @@ def __init__( ak="", sk="", sts_token="", + api_key="", scheme="https", connection_timeout=30, socket_timeout=30, @@ -89,6 +90,12 @@ def __init__( raise ValueError( "DATABASE_VIKINGMEM_BASE_URL must start with http:// or https://" ) + if (api_key or "").strip(): + raise ValueError( + "VikingDB memory collection management requires AK/SK or IAM " + "credentials. API key auth is only supported by the VikingMem SDK " + "for memory operations on existing collections." + ) self.service_info = VikingDBMemoryClient.get_service_info( host, region, scheme, connection_timeout, socket_timeout diff --git a/veadk/knowledgebase/backends/vikingdb_knowledge_backend.py b/veadk/knowledgebase/backends/vikingdb_knowledge_backend.py index e5c6e5508..ce3ff073a 100644 --- a/veadk/knowledgebase/backends/vikingdb_knowledge_backend.py +++ b/veadk/knowledgebase/backends/vikingdb_knowledge_backend.py @@ -68,6 +68,17 @@ def _viking_session_token_from_env() -> str: return os.getenv("VOLCENGINE_SESSION_TOKEN", "") +def _clean_api_key(value: str | None) -> str | None: + value = (value or "").strip() + if not value or value.lower() in {"none", "null"}: + return None + return value + + +def _viking_api_key_from_env() -> str | None: + return _clean_api_key(os.getenv("DATABASE_VIKING_API_KEY")) + + def _byteplus_viking_region(region: str | None) -> str: """Return the supported BytePlus VikingDB Knowledge Base region.""" region = (region or "").strip() @@ -194,6 +205,7 @@ class VikingDBKnowledgeBackend(BaseKnowledgebaseBackend): default_factory=_viking_secret_key_from_env ) session_token: str = Field(default_factory=_viking_session_token_from_env) + api_key: str | None = Field(default_factory=_viking_api_key_from_env) volcengine_project: str = Field( default_factory=lambda: os.getenv("DATABASE_VIKING_PROJECT", "default") @@ -206,9 +218,7 @@ class VikingDBKnowledgeBackend(BaseKnowledgebaseBackend): default_factory=lambda: os.getenv("DATABASE_VIKING_VERSION", "2") ) - cloud_provider: str = Field( - default_factory=lambda: os.getenv("CLOUD_PROVIDER", "volces") - ) + cloud_provider: str = Field(default_factory=_viking_cloud_provider) region: str = Field(default="") base_url: str = Field(default="") @@ -220,6 +230,7 @@ class VikingDBKnowledgeBackend(BaseKnowledgebaseBackend): _viking_sdk_client = None def model_post_init(self, __context: Any) -> None: + self.api_key = _clean_api_key(self.api_key) or _viking_api_key_from_env() if self.cloud_provider.lower() == "byteplus": self.region = _byteplus_viking_region( self.region or os.getenv("DATABASE_VIKING_REGION") @@ -238,21 +249,39 @@ def model_post_init(self, __context: Any) -> None: ): self.tos_config.region = self.region self.tos_config.endpoint = f"tos-{self.region}.bytepluses.com" - elif not self.region: + else: self.region = ( - os.getenv("DATABASE_VIKING_REGION") + self.region + or os.getenv("DATABASE_VIKING_REGION") or os.getenv("REGION") or "cn-beijing" ) - self.base_url = f"https://api-knowledgebase.mlp.{self.region}.volces.com" - self.host = f"api-knowledgebase.mlp.{self.region}.volces.com" + self.base_url = ( + self.base_url + or f"https://api-knowledgebase.mlp.{self.region}.volces.com" + ) + self.host = self.host or f"api-knowledgebase.mlp.{self.region}.volces.com" logger.info(f"Cloud provider: {self.cloud_provider.lower()}") logger.info(f"VikingDBKnowledgeBackend: region={self.region}, host={self.host}") + logger.info( + "VikingDBKnowledgeBackend auth: " + + ( + "API key for knowledge search; AK/SK or IAM for management" + if self.api_key + else "AK/SK or IAM" + ) + ) self.precheck_index_naming() # check whether collection exist, if not, create it + if self.api_key and not self._has_explicit_management_credentials(): + logger.info( + "Skip VikingDB knowledgebase collection management precheck: " + "API key is configured, but AK/SK credentials are not configured." + ) + return if not self.collection_status()["existed"]: logger.warning( f"VikingDB knowledgebase collection {self.index} does not exist, please create it first..." @@ -270,15 +299,15 @@ def precheck_index_naming(self): "it must start with an English letter, contain only letters, numbers, and underscores, and have a length of 1-128." ) + def _has_explicit_management_credentials(self) -> bool: + return bool(self.volcengine_access_key and self.volcengine_secret_key) + def _get_tos_client(self, tos_bucket_name: str) -> VeTOS: ak = None sk = None sts_token = None if not (self.volcengine_access_key and self.volcengine_secret_key): - cred = self._set_service_info() - ak = cred.access_key_id - sk = cred.secret_access_key - sts_token = cred.session_token + ak, sk, sts_token = self._get_ak_sk_sts() return VeTOS( ak=ak or self.volcengine_access_key, @@ -669,32 +698,54 @@ def _search_knowledge( "chunk_diffusion_count": chunk_diffusion_count, } - ak = None - sk = None - sts_token = None - if not (self.volcengine_access_key and self.volcengine_secret_key): - cred = self._set_service_info() - ak = cred.access_key_id - sk = cred.secret_access_key - sts_token = cred.session_token + if self.api_key: + logger.info( + "Search VikingDB knowledgebase using API key auth: " + f"collection={self.index}, project={self.volcengine_project}" + ) + body = { + "name": self.index, + "project": self.volcengine_project, + "query": query, + "limit": top_k, + "dense_weight": 0.5, + "post_processing": post_precessing, + } + if query_param is not None: + body["query_param"] = query_param + response = self._do_request( + body=self._with_resource_id(body), + path="/api/knowledge/collection/search_knowledge", + method="POST", + auth_mode="api_key", + ) + if response.get("code") not in (0, None): + raise ValueError(f"Error during knowledge search: {response}") + response = response.get("data", response) + else: + ak, sk, sts_token = self._get_ak_sk_sts() + logger.info( + "Search VikingDB knowledgebase using AK/SK or IAM auth: " + f"collection={self.index}, project={self.volcengine_project}" + ) - self._viking_sdk_client = VikingKnowledgeBaseService( - host=self.host, - ak=ak or self.volcengine_access_key, - sk=sk or self.volcengine_secret_key, - sts_token=sts_token or self.session_token, - scheme=self.schema, - ) + self._viking_sdk_client = VikingKnowledgeBaseService( + host=self.host, + ak=ak, + sk=sk, + sts_token=sts_token, + scheme=self.schema, + ) - response = self._viking_sdk_client.search_knowledge( - collection_name=self.index, - project=self.volcengine_project, - query=query, - limit=top_k, - query_param=query_param, - post_processing=post_precessing, - resource_id=self.resource_id or None, - ) + response = self._viking_sdk_client.search_knowledge( + collection_name=self.index, + project=self.volcengine_project, + query=query, + limit=top_k, + query_param=query_param, + post_processing=post_precessing, + resource_id=self.resource_id or None, + ) logger.debug( f"Search knowledge {self.index} using project {self.volcengine_project} original response: {response}" @@ -744,38 +795,62 @@ def _set_service_info(self) -> VeIAMCredential: cred = get_credential_from_vefaas_iam() return cred + def _get_ak_sk_sts(self) -> tuple[str, str, str]: + if self.volcengine_access_key and self.volcengine_secret_key: + return ( + self.volcengine_access_key, + self.volcengine_secret_key, + self.session_token, + ) + cred = self._set_service_info() + return cred.access_key_id, cred.secret_access_key, cred.session_token + def _do_request( self, body: dict, path: str, method: Literal["GET", "POST", "PUT", "DELETE"] = "POST", + auth_mode: Literal["ak_sk", "api_key"] = "ak_sk", ) -> dict: full_path = f"{self.base_url}{path}" - ak = None - sk = None - sts_token = None - if not (self.volcengine_access_key and self.volcengine_secret_key): - cred = self._set_service_info() - ak = cred.access_key_id - sk = cred.secret_access_key - sts_token = cred.session_token - - request = build_vikingdb_knowledgebase_request( - path=path, - volcengine_access_key=ak or self.volcengine_access_key, - volcengine_secret_key=sk or self.volcengine_secret_key, - session_token=sts_token or self.session_token, - method=method, - data=body, - region=self.region, - ) - response = requests.request( - method=method, - url=full_path, - headers=request.headers, - data=request.body, - ) + if auth_mode == "api_key": + if not self.api_key: + raise ValueError("VikingDB API key is required for API key auth mode.") + logger.debug( + "VikingDB knowledgebase request uses API key auth: path=%s", path + ) + response = requests.request( + method=method, + url=full_path, + headers={ + "Accept": "application/json", + "Content-Type": "application/json", + "Authorization": f"Bearer {self.api_key}", + }, + data=json.dumps(body), + ) + else: + ak, sk, sts_token = self._get_ak_sk_sts() + logger.debug( + "VikingDB knowledgebase request uses AK/SK or IAM auth: path=%s", + path, + ) + request = build_vikingdb_knowledgebase_request( + path=path, + volcengine_access_key=ak, + volcengine_secret_key=sk, + session_token=sts_token, + method=method, + data=body, + region=self.region, + ) + response = requests.request( + method=method, + url=full_path, + headers=request.headers, + data=request.body, + ) if not response.ok: logger.error( f"VikingDBKnowledgeBackend error during request: {response.json()}" diff --git a/veadk/memory/long_term_memory_backends/vikingdb_memory_backend.py b/veadk/memory/long_term_memory_backends/vikingdb_memory_backend.py index 090c36cd3..896c19407 100644 --- a/veadk/memory/long_term_memory_backends/vikingdb_memory_backend.py +++ b/veadk/memory/long_term_memory_backends/vikingdb_memory_backend.py @@ -21,7 +21,7 @@ from pydantic import Field from typing_extensions import override -from vikingdb import IAM +from vikingdb import IAM, APIKey from vikingdb.memory import VikingMem import veadk.config # noqa E401 @@ -37,6 +37,7 @@ DEFAULT_BYTEPLUS_VIKING_MEMORY_REGION, ) from veadk.utils.logger import get_logger +from veadk.utils.misc import getenv logger = get_logger(__name__) @@ -67,6 +68,17 @@ def _viking_session_token_from_env() -> str: return os.getenv("VOLCENGINE_SESSION_TOKEN", "") +def _clean_api_key(value: str | None) -> str | None: + value = (value or "").strip() + if not value or value.lower() in {"none", "null"}: + return None + return value + + +def _vikingmem_api_key_from_env() -> str | None: + return _clean_api_key(os.getenv("DATABASE_VIKINGMEM_API_KEY")) + + class VikingDBLTMBackend(BaseLongTermMemoryBackend): volcengine_access_key: str | None = Field( default_factory=_viking_access_key_from_env @@ -77,10 +89,9 @@ class VikingDBLTMBackend(BaseLongTermMemoryBackend): ) session_token: str = Field(default_factory=_viking_session_token_from_env) + api_key: str | None = Field(default_factory=_vikingmem_api_key_from_env) - cloud_provider: str = Field( - default_factory=lambda: os.getenv("CLOUD_PROVIDER", "volces") - ) + cloud_provider: str = Field(default_factory=_viking_cloud_provider) region: str = Field(default="") """VikingDB memory region""" @@ -93,6 +104,7 @@ class VikingDBLTMBackend(BaseLongTermMemoryBackend): memory_type: list[str] = Field(default_factory=list) def model_post_init(self, __context: Any, /) -> None: + self.api_key = _clean_api_key(self.api_key) or _vikingmem_api_key_from_env() if self.cloud_provider.lower() == "byteplus": self.region = DEFAULT_BYTEPLUS_VIKING_MEMORY_REGION elif not self.region: @@ -116,8 +128,22 @@ def model_post_init(self, __context: Any, /) -> None: self.memory_type = ["sys_event_v1", "sys_profile_v1"] logger.info(f"Using memory type: {self.memory_type}") + logger.info( + "VikingDBLTMBackend auth: " + + ( + "API key for memory operations; AK/SK or IAM for collection management" + if self.api_key + else "AK/SK or IAM" + ) + ) # check whether collection exist, if not, create it + if self.api_key and not self._has_explicit_management_credentials(): + logger.info( + "Skip VikingDB memory collection management precheck: " + "API key is configured, but AK/SK credentials are not configured." + ) + return if not self._collection_exist(): self._create_collection() @@ -158,6 +184,9 @@ def _create_collection(self) -> None: logger.debug(f"Create collection with response {response}") return response + def _has_explicit_management_credentials(self) -> bool: + return bool(self.volcengine_access_key and self.volcengine_secret_key) + def _get_ak_sk_sts(self) -> tuple[str, str, str]: ak = "" sk = "" @@ -182,14 +211,38 @@ def _get_ak_sk_sts(self) -> tuple[str, str, str]: return ak, sk, sts_token + def _get_viking_memory_host(self) -> str: + if self.cloud_provider.lower() == "byteplus": + return DEFAULT_BYTEPLUS_VIKING_MEMORY_HOST + return f"api-knowledgebase.mlp.{self.region}.volces.com" + + def _get_viking_memory_endpoint(self) -> tuple[str, str]: + host = self._get_viking_memory_host() + scheme = "https" + env_host = getenv( + "DATABASE_VIKINGMEM_BASE_URL", + default_value=None, + allow_false_values=True, + ) + if env_host: + if env_host.startswith("http://"): + host = env_host.replace("http://", "") + scheme = "http" + elif env_host.startswith("https://"): + host = env_host.replace("https://", "") + scheme = "https" + else: + raise ValueError( + "DATABASE_VIKINGMEM_BASE_URL must start with http:// or https://" + ) + return host, scheme + def _get_client(self) -> VikingDBMemoryClient: ak, sk, sts_token = self._get_ak_sk_sts() - if self.cloud_provider.lower() == "byteplus": - host = DEFAULT_BYTEPLUS_VIKING_MEMORY_HOST - else: - host = f"api-knowledgebase.mlp.{self.region}.volces.com" + host, scheme = self._get_viking_memory_endpoint() logger.info(f"Cloud provider: {self.cloud_provider.lower()}") logger.info(f"VikingDBLTMBackend: region={self.region}, host={host}") + logger.info("VikingDB memory collection management uses AK/SK or IAM auth") return VikingDBMemoryClient( host=host, @@ -197,33 +250,28 @@ def _get_client(self) -> VikingDBMemoryClient: sk=sk, sts_token=sts_token, region=self.region, + scheme=scheme, ) def _get_sdk_client(self) -> VikingMem: - ak, sk, sts_token = self._get_ak_sk_sts() - if self.cloud_provider.lower() == "byteplus": - host = DEFAULT_BYTEPLUS_VIKING_MEMORY_HOST - else: - host = f"api-knowledgebase.mlp.{self.region}.volces.com" + host, scheme = self._get_viking_memory_endpoint() logger.info(f"Cloud provider: {self.cloud_provider.lower()}") logger.info(f"VikingDBLTMBackend: region={self.region}, host={host}") - client = VikingDBMemoryClient( - host=host, - region=self.region, - ak=ak, - sk=sk, - sts_token=sts_token, - ) - + sts_token = "" + if self.api_key: + logger.info("VikingDB memory SDK uses API key auth") + auth = APIKey(api_key=self.api_key) + else: + ak, sk, sts_token = self._get_ak_sk_sts() + logger.info("VikingDB memory SDK uses AK/SK or IAM auth") + auth = IAM(ak=ak, sk=sk) return VikingMem( - host=client.get_host(), + host=host, region=self.region, - auth=IAM( - ak=ak, - sk=sk, - ), + auth=auth, sts_token=sts_token, + scheme=scheme, ) @override @@ -311,44 +359,34 @@ def search_memory( ) def get_user_profile(self, user_id: str) -> str: - from veadk.utils.volcengine_sign import ve_request - - response: dict = self._get_client().get_collection( - collection_name=self.index, project=self.volcengine_project - ) - - mem_id = response["Result"]["ResourceId"] logger.info( - f"Get user profile for user_id={user_id} from Viking Memory with mem_id={mem_id}" + f"Get user profile for user_id={user_id} from Viking Memory collection={self.index}" ) - - ak, sk, sts_token = self._get_ak_sk_sts() - response = ve_request( - request_body={ - "filter": { - "user_id": [user_id], - "memory_category": 1, - }, - "limit": 5000, - "resource_id": mem_id, + client = self._get_sdk_client() + collection = client.get_collection( + collection_name=self.index, project_name=self.volcengine_project + ) + response = collection.search_memory( + filter={ + "user_id": [user_id], + "memory_category": 1, }, - action="MemorySearch", - ak=ak, - sk=sk, - header={"X-Security-Token": sts_token}, - service="vikingdb", - version="2025-06-09", - region=self.region, - host="open.volcengineapi.com", + limit=5000, ) - try: + code = response.get("code") + if code is not None and code != 0: + raise ValueError(f"Get VikingDB user profile error: {response}") + + result_list = response.get("data", {}).get("result_list", []) + if not result_list: logger.debug( - f"Response from VikingDB: {response}, user_profile: {response['data']['result_list'][0]['memory_info']['user_profile']}" - ) - return response["data"]["result_list"][0]["memory_info"]["user_profile"] - except (KeyError, IndexError): - logger.error( - f"Failed to get user profile for user_id={user_id} mem_id={mem_id}: {response}" + f"No VikingDB user profile found for user_id={user_id} collection={self.index}: {response}" ) return "" + + user_profile = result_list[0].get("memory_info", {}).get("user_profile", "") + logger.debug( + f"Response from VikingDB: {response}, user_profile: {user_profile}" + ) + return user_profile