Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,45 @@
|
||||
"""
|
||||
向量存储模块
|
||||
|
||||
提供可插拔的向量存储后端,支持 Qdrant 等专业向量数据库。
|
||||
与业务数据库完全解耦,segment 表只存业务数据,向量数据存在向量数据库中。
|
||||
"""
|
||||
from ai_platform.knowledge.vector_store.base import BaseVectorStore, VectorPoint, VectorSearchResult
|
||||
from ai_platform.knowledge.vector_store.qdrant_store import QdrantVectorStore
|
||||
|
||||
__all__ = [
|
||||
'BaseVectorStore',
|
||||
'VectorPoint',
|
||||
'VectorSearchResult',
|
||||
'QdrantVectorStore',
|
||||
'get_vector_store',
|
||||
]
|
||||
|
||||
# 单例缓存
|
||||
_vector_store_instance: BaseVectorStore | None = None
|
||||
|
||||
|
||||
def get_vector_store() -> BaseVectorStore:
|
||||
"""
|
||||
工厂函数:根据配置获取向量存储实例(单例)
|
||||
"""
|
||||
global _vector_store_instance
|
||||
if _vector_store_instance is not None:
|
||||
return _vector_store_instance
|
||||
|
||||
from app.config import settings
|
||||
|
||||
store_type = getattr(settings, 'VECTOR_STORE_TYPE', 'qdrant')
|
||||
|
||||
if store_type == 'qdrant':
|
||||
_vector_store_instance = QdrantVectorStore(
|
||||
host=getattr(settings, 'QDRANT_HOST', 'localhost'),
|
||||
port=getattr(settings, 'QDRANT_PORT', 6333),
|
||||
api_key=getattr(settings, 'QDRANT_API_KEY', None),
|
||||
grpc_port=getattr(settings, 'QDRANT_GRPC_PORT', 6334),
|
||||
prefer_grpc=getattr(settings, 'QDRANT_PREFER_GRPC', False),
|
||||
)
|
||||
else:
|
||||
raise ValueError(f'不支持的向量存储类型: {store_type}')
|
||||
|
||||
return _vector_store_instance
|
||||
@@ -0,0 +1,138 @@
|
||||
"""
|
||||
向量存储抽象基类
|
||||
|
||||
定义向量存储的统一接口,所有向量存储后端必须实现这些方法。
|
||||
"""
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class VectorPoint:
|
||||
"""向量数据点"""
|
||||
id: str
|
||||
vector: List[float]
|
||||
payload: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class VectorSearchResult:
|
||||
"""向量搜索结果"""
|
||||
id: str
|
||||
score: float
|
||||
payload: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
class BaseVectorStore(ABC):
|
||||
"""
|
||||
向量存储抽象基类
|
||||
|
||||
每个知识库对应一个 collection,collection 名称格式: kb_{knowledge_base_id}
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def collection_name(knowledge_base_id: str) -> str:
|
||||
"""生成 collection 名称"""
|
||||
return f"kb_{knowledge_base_id.replace('-', '_')}"
|
||||
|
||||
@abstractmethod
|
||||
async def ensure_collection(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
vector_size: int,
|
||||
) -> None:
|
||||
"""
|
||||
确保 collection 存在,不存在则创建
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
vector_size: 向量维度
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete_collection(self, knowledge_base_id: str) -> None:
|
||||
"""
|
||||
删除 collection(删除知识库时调用)
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def upsert(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
points: List[VectorPoint],
|
||||
) -> None:
|
||||
"""
|
||||
批量写入/更新向量
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
points: 向量数据点列表
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
point_ids: List[str],
|
||||
) -> None:
|
||||
"""
|
||||
批量删除向量
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
point_ids: 要删除的向量 ID 列表
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def search(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
query_vector: List[float],
|
||||
top_k: int = 10,
|
||||
score_threshold: float = 0.0,
|
||||
filter_conditions: Optional[Dict[str, Any]] = None,
|
||||
) -> List[VectorSearchResult]:
|
||||
"""
|
||||
向量相似度搜索
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
query_vector: 查询向量
|
||||
top_k: 返回数量
|
||||
score_threshold: 最低相似度阈值
|
||||
filter_conditions: 过滤条件(如 {"document_id": "xxx"})
|
||||
|
||||
Returns:
|
||||
搜索结果列表,按相似度降序排列
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete_by_filter(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
filter_conditions: Dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
按条件删除向量(如删除某个文档的所有向量)
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
filter_conditions: 过滤条件(如 {"document_id": "xxx"})
|
||||
"""
|
||||
...
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""健康检查"""
|
||||
return True
|
||||
@@ -0,0 +1,260 @@
|
||||
"""
|
||||
Qdrant 向量存储实现
|
||||
|
||||
使用 Qdrant 作为向量数据库后端,通过 qdrant-client 进行交互。
|
||||
每个知识库对应一个 Qdrant collection。
|
||||
"""
|
||||
import logging
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from qdrant_client import AsyncQdrantClient
|
||||
from qdrant_client.models import (
|
||||
Distance,
|
||||
FieldCondition,
|
||||
Filter,
|
||||
FilterSelector,
|
||||
MatchValue,
|
||||
PointIdsList,
|
||||
PointStruct,
|
||||
VectorParams,
|
||||
)
|
||||
|
||||
from ai_platform.knowledge.vector_store.base import (
|
||||
BaseVectorStore,
|
||||
VectorPoint,
|
||||
VectorSearchResult,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class QdrantVectorStore(BaseVectorStore):
|
||||
"""
|
||||
Qdrant 向量存储
|
||||
|
||||
特性:
|
||||
- 高性能向量检索(HNSW 索引)
|
||||
- 支持 payload 过滤
|
||||
- 支持 REST 和 gRPC 协议
|
||||
- 与业务数据库完全解耦
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
host: str = "localhost",
|
||||
port: int = 6333,
|
||||
api_key: Optional[str] = None,
|
||||
grpc_port: int = 6334,
|
||||
prefer_grpc: bool = False,
|
||||
):
|
||||
self._host = host
|
||||
self._port = port
|
||||
self._api_key = api_key
|
||||
self._grpc_port = grpc_port
|
||||
self._prefer_grpc = prefer_grpc
|
||||
self._client: Optional[AsyncQdrantClient] = None
|
||||
|
||||
async def _get_client(self) -> AsyncQdrantClient:
|
||||
"""获取或创建 Qdrant 客户端(懒初始化)"""
|
||||
if self._client is None:
|
||||
# 如果 host 已包含协议前缀,直接作为 url 使用;
|
||||
# 否则拼接 http:// 避免 qdrant-client 对非 localhost 域名自动走 HTTPS
|
||||
if self._host.startswith("http://") or self._host.startswith("https://"):
|
||||
url = f"{self._host}:{self._port}"
|
||||
else:
|
||||
url = f"http://{self._host}:{self._port}"
|
||||
self._client = AsyncQdrantClient(
|
||||
url=url,
|
||||
api_key=self._api_key,
|
||||
grpc_port=self._grpc_port,
|
||||
prefer_grpc=self._prefer_grpc,
|
||||
timeout=30,
|
||||
)
|
||||
return self._client
|
||||
|
||||
async def ensure_collection(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
vector_size: int,
|
||||
) -> None:
|
||||
"""确保 collection 存在且维度匹配"""
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
|
||||
collections = await client.get_collections()
|
||||
existing_names = {c.name for c in collections.collections}
|
||||
|
||||
if name in existing_names:
|
||||
# 检查已有 collection 的维度是否匹配
|
||||
info = await client.get_collection(collection_name=name)
|
||||
existing_size = info.config.params.vectors.size
|
||||
if existing_size != vector_size:
|
||||
logger.warning(
|
||||
f"Qdrant collection {name} 维度不匹配: "
|
||||
f"已有={existing_size}, 期望={vector_size},删除重建"
|
||||
)
|
||||
await client.delete_collection(collection_name=name)
|
||||
else:
|
||||
return
|
||||
|
||||
await client.create_collection(
|
||||
collection_name=name,
|
||||
vectors_config=VectorParams(
|
||||
size=vector_size,
|
||||
distance=Distance.COSINE,
|
||||
),
|
||||
)
|
||||
# 创建 payload 索引,加速过滤查询
|
||||
await client.create_payload_index(
|
||||
collection_name=name,
|
||||
field_name="document_id",
|
||||
field_schema="keyword",
|
||||
)
|
||||
logger.info(f"Qdrant collection 已创建: {name} (dim={vector_size})")
|
||||
|
||||
async def delete_collection(self, knowledge_base_id: str) -> None:
|
||||
"""删除 collection"""
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
try:
|
||||
await client.delete_collection(collection_name=name)
|
||||
logger.info(f"Qdrant collection 已删除: {name}")
|
||||
except Exception as e:
|
||||
logger.warning(f"删除 Qdrant collection 失败: {name}, {e}")
|
||||
|
||||
async def upsert(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
points: List[VectorPoint],
|
||||
) -> None:
|
||||
"""批量写入/更新向量"""
|
||||
if not points:
|
||||
return
|
||||
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
|
||||
qdrant_points = [
|
||||
PointStruct(
|
||||
id=self._to_uuid(p.id),
|
||||
vector=p.vector,
|
||||
payload={**p.payload, 'segment_id': p.id},
|
||||
)
|
||||
for p in points
|
||||
]
|
||||
|
||||
# Qdrant 单次 upsert 建议不超过 100 个点
|
||||
batch_size = 100
|
||||
for i in range(0, len(qdrant_points), batch_size):
|
||||
batch = qdrant_points[i:i + batch_size]
|
||||
await client.upsert(
|
||||
collection_name=name,
|
||||
points=batch,
|
||||
)
|
||||
|
||||
async def delete(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
point_ids: List[str],
|
||||
) -> None:
|
||||
"""批量删除向量"""
|
||||
if not point_ids:
|
||||
return
|
||||
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
|
||||
uuid_ids = [self._to_uuid(pid) for pid in point_ids]
|
||||
await client.delete(
|
||||
collection_name=name,
|
||||
points_selector=PointIdsList(points=uuid_ids),
|
||||
)
|
||||
|
||||
async def search(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
query_vector: List[float],
|
||||
top_k: int = 10,
|
||||
score_threshold: float = 0.0,
|
||||
filter_conditions: Optional[Dict[str, Any]] = None,
|
||||
) -> List[VectorSearchResult]:
|
||||
"""向量相似度搜索"""
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
|
||||
# 构建过滤条件
|
||||
query_filter = self._build_filter(filter_conditions) if filter_conditions else None
|
||||
|
||||
try:
|
||||
results = await client.search(
|
||||
collection_name=name,
|
||||
query_vector=query_vector,
|
||||
limit=top_k,
|
||||
score_threshold=score_threshold,
|
||||
query_filter=query_filter,
|
||||
with_payload=True,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Qdrant 搜索失败: {e}")
|
||||
return []
|
||||
|
||||
return [
|
||||
VectorSearchResult(
|
||||
id=(hit.payload or {}).get('segment_id', str(hit.id)),
|
||||
score=hit.score,
|
||||
payload=hit.payload or {},
|
||||
)
|
||||
for hit in results
|
||||
]
|
||||
|
||||
async def delete_by_filter(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
filter_conditions: Dict[str, Any],
|
||||
) -> None:
|
||||
"""按条件删除向量"""
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
|
||||
query_filter = self._build_filter(filter_conditions)
|
||||
if query_filter:
|
||||
await client.delete(
|
||||
collection_name=name,
|
||||
points_selector=FilterSelector(filter=query_filter),
|
||||
)
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""健康检查"""
|
||||
try:
|
||||
client = await self._get_client()
|
||||
await client.get_collections()
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Qdrant 健康检查失败: {e}")
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _to_uuid(string_id: str) -> str:
|
||||
"""将任意字符串 ID 确定性转换为 UUID5(Qdrant 要求 point ID 为 UUID 或整数)"""
|
||||
return str(uuid.uuid5(uuid.NAMESPACE_DNS, string_id))
|
||||
|
||||
@staticmethod
|
||||
def _build_filter(conditions: Dict[str, Any]) -> Optional[Filter]:
|
||||
"""构建 Qdrant 过滤条件"""
|
||||
if not conditions:
|
||||
return None
|
||||
|
||||
must = []
|
||||
for key, value in conditions.items():
|
||||
if isinstance(value, list):
|
||||
# 列表值:任一匹配(OR 语义),用 should 包裹后作为一个 must 条件
|
||||
should_conditions = [
|
||||
FieldCondition(key=key, match=MatchValue(value=v))
|
||||
for v in value
|
||||
]
|
||||
must.append(Filter(should=should_conditions))
|
||||
else:
|
||||
must.append(FieldCondition(key=key, match=MatchValue(value=value)))
|
||||
|
||||
return Filter(must=must) if must else None
|
||||
Reference in New Issue
Block a user