- Updated test files to import `AgentSessionLocal` from `advisor_db` instead of directly, preventing session binding to the real database during tests. - Fixed 6 test cases to use the new login token utility, ensuring consistency across authentication methods. - Adjusted customer risk codes in `test_convert_confirm.py` to reflect changes in customer classification (C3 to C4). - Verified that changes resulted in zero database pollution during test runs, maintaining integrity of the testing environment. - Documented findings and updates in the relevant test logs and memory files, ensuring clarity on the current state of tests and defects.
207 lines
7.5 KiB
Python
207 lines
7.5 KiB
Python
import subprocess
|
|
import sys
|
|
from pathlib import Path
|
|
from uuid import uuid4
|
|
|
|
from fastapi.testclient import TestClient
|
|
|
|
from app import advisor_db
|
|
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:
|
|
from tests.advisor_test_utils import login_token
|
|
|
|
return login_token(username, password)
|
|
|
|
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 advisor_db.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 advisor_db.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
|