Files
group_xinghuo_jinrong/tests/test_sprint2_template_vectors.py
T
zhanghongyu_0626 70aa861983 feat(advisor-agent): Introduce advisor agent functionalities with compliance, KYC, and script templates
- Added new modules for advisor compliance, KYC sessions, and script templates, enhancing the advisor agent's capabilities.
- Implemented a comprehensive API structure under the `/api/advisor-agent` prefix, ensuring clear organization and access to new features.
- Established database models and repositories for compliance rules and KYC sessions, facilitating robust data management.
- Integrated exception handling and response models to improve error management and user feedback.
- Updated settings to include new configurations for compliance and KYC features, ensuring flexibility and adaptability.

This update significantly expands the advisor agent's functionality, providing essential tools for compliance and customer interaction while maintaining a structured API design.
2026-09-12 16:33:07 +08:00

210 lines
7.6 KiB
Python

import subprocess
import sys
from pathlib import Path
from uuid import uuid4
from fastapi.testclient import TestClient
from app.advisor_db import AgentSessionLocal
from app.main import app
from app.model.entities_advisor import ScriptTemplate
from app.model.advisor_schemas import AuthContext, TemplateUpdate
from app.repository.script_template_repository import ScriptTemplateRepository
from app.service.script_template_service import ScriptTemplateService
ROOT = Path(__file__).resolve().parents[1]
SCRIPT = ROOT / "scripts" / "sync" / "sync_template_vectors.py"
client = TestClient(app)
class FakeEmbeddingTool:
def embed_text(self, text: str) -> list[float]:
assert text
return [0.125] * 1024
class FakeVectorStore:
def __init__(self) -> None:
self.ensured_collections: list[str] = []
self.upserted: dict[str, dict] = {}
self.deleted_template_ids: list[int] = []
def ensure_template_collection(self) -> None:
self.ensured_collections.append("kb_script_templates")
def upsert_template(self, record) -> str:
assert len(record.embedding) == 1024
self.upserted[record.vector_id] = {
"template_id": record.template_id,
"scene": record.scene,
"title": record.title,
"tags": record.tags,
"chunk_text": record.chunk_text,
"chunk_no": record.chunk_no,
}
return record.vector_id
def delete_template(self, template_id: int) -> None:
self.deleted_template_ids.append(template_id)
def token_for(username: str, password: str) -> str:
response = client.post(
"/api/v1/auth/login",
json={"username": username, "password": password},
)
assert response.status_code == 200
return response.json()["data"]["access_token"]
def compliance_auth() -> AuthContext:
token = token_for("compliance_test", "compliance_test")
from app.service.auth_service import auth_service
return auth_service.parse_token(token, f"trace-vector-auth-{uuid4().hex}")
def create_template_row(
*,
is_approved: bool,
is_active: bool = True,
created_by: str = "test:template_vectors",
) -> ScriptTemplate:
suffix = uuid4().hex[:8]
with AgentSessionLocal() as session:
template = ScriptTemplate(
scene="loss_comfort",
customer_type="C3",
title=f"向量同步模板 {suffix}",
content=f"您好,近期市场波动较大,请结合风险承受能力理性看待。{suffix}",
tags=["亏损", "安抚"],
is_approved=is_approved,
approved_by="compliance_test" if is_approved else None,
is_active=is_active,
version=1,
usage_count=0,
created_by=created_by,
updated_by=created_by,
)
session.add(template)
session.commit()
session.refresh(template)
session.expunge(template)
return template
def load_embedding_id(template_id: int) -> str | None:
with AgentSessionLocal() as session:
template = session.get(ScriptTemplate, template_id)
assert template is not None
return template.embedding_id
def test_template_vector_sync_indexes_only_approved_active_templates_and_updates_embedding_id():
from app.service.script_template_vector_service import ScriptTemplateVectorService
created_by = f"test:template_vectors:{uuid4().hex}"
approved = create_template_row(is_approved=True, created_by=created_by)
unapproved = create_template_row(is_approved=False, created_by=created_by)
inactive = create_template_row(is_approved=True, is_active=False, created_by=created_by)
vector_store = FakeVectorStore()
service = ScriptTemplateVectorService(
repository=ScriptTemplateRepository(),
embedding_tool=FakeEmbeddingTool(),
vector_store=vector_store,
)
result = service.sync_approved_templates(created_by=created_by)
assert result.total == 1
assert result.upserted == 1
assert result.skipped == 0
assert f"tpl_{approved.id}_0" in vector_store.upserted
assert f"tpl_{unapproved.id}_0" not in vector_store.upserted
assert f"tpl_{inactive.id}_0" not in vector_store.upserted
assert load_embedding_id(approved.id) == f"tpl_{approved.id}_0"
def test_template_service_approval_upserts_vector_and_content_update_deletes_vector():
from app.service.script_template_vector_service import ScriptTemplateVectorService
template = create_template_row(is_approved=False, created_by=f"test:template_vectors:{uuid4().hex}")
vector_store = FakeVectorStore()
vector_service = ScriptTemplateVectorService(
repository=ScriptTemplateRepository(),
embedding_tool=FakeEmbeddingTool(),
vector_store=vector_store,
)
template_service = ScriptTemplateService(repository=ScriptTemplateRepository(), vector_service=vector_service)
auth = compliance_auth()
approved = template_service.update_template(template.id, TemplateUpdate(is_approved=True), auth)
updated = template_service.update_template(
template.id,
TemplateUpdate(content="您好,市场短期波动较大,请先阅读风险揭示材料。"),
auth,
)
assert approved.embedding_id == f"tpl_{template.id}_0"
assert f"tpl_{template.id}_0" in vector_store.upserted
assert updated.is_approved is False
assert updated.embedding_id is None
assert template.id in vector_store.deleted_template_ids
def test_template_vector_service_rejects_wrong_embedding_dimension():
from app.service.script_template_vector_service import (
TemplateVectorError,
ScriptTemplateVectorService,
)
class BadEmbeddingTool:
def embed_text(self, text: str) -> list[float]:
return [0.1, 0.2]
template = create_template_row(is_approved=True, created_by=f"test:template_vectors:{uuid4().hex}")
service = ScriptTemplateVectorService(
repository=ScriptTemplateRepository(),
embedding_tool=BadEmbeddingTool(),
vector_store=FakeVectorStore(),
)
try:
service.upsert_template(template)
except TemplateVectorError as exc:
assert "1024" in str(exc)
else:
raise AssertionError("TemplateVectorError was not raised")
def test_milvus_loader_ignores_relative_env_file_uri():
from app.tool.milvus_tool import _load_pymilvus
MilvusClient, DataType = _load_pymilvus()
assert MilvusClient.__name__ == "MilvusClient"
assert hasattr(DataType, "FLOAT_VECTOR")
def test_template_vector_sync_skips_when_embedding_backend_is_unavailable():
from app.service.script_template_vector_service import ScriptTemplateVectorService
from app.tool.embedding_tool import EmbeddingError
class FailingEmbeddingTool:
def embed_text(self, text: str) -> list[float]:
raise EmbeddingError("embedding backend unavailable")
created_by = f"test:template_vectors:{uuid4().hex}"
create_template_row(is_approved=True, created_by=created_by)
service = ScriptTemplateVectorService(
repository=ScriptTemplateRepository(),
embedding_tool=FailingEmbeddingTool(),
vector_store=FakeVectorStore(),
)
result = service.sync_approved_templates(created_by=created_by)
assert result.total >= 1
assert result.skipped >= 1
def test_template_vector_sync_script_supports_dry_run():
completed = subprocess.run(
[sys.executable, str(SCRIPT), "--dry-run"],
check=True,
capture_output=True,
text=True,
cwd=ROOT,
)
assert "Template vector sync dry run:" in completed.stdout
assert "total=" in completed.stdout