feat:新增投顾agent和nl2sqlagent

This commit is contained in:
2026-09-13 16:19:24 +08:00
parent c80c6acac0
commit 163192bf55
122 changed files with 7488 additions and 362 deletions
+86
View File
@@ -0,0 +1,86 @@
"""NL2SQL 元数据向量的 Milvus 数据库和集合初始化。"""
from __future__ import annotations
from pymilvus import AsyncMilvusClient, DataType
from config.settings import settings
NL2SQL_DATABASE_NAME = settings.milvus.db_name or "mutual_fund"
NL2SQL_COLLECTION = "data_agent_chunks"
def build_nl2sql_schema():
schema = AsyncMilvusClient.create_schema(auto_id=False, enable_dynamic_field=False)
schema.add_field("id", DataType.VARCHAR, is_primary=True, max_length=128)
schema.add_field("vector", DataType.FLOAT_VECTOR, dim=settings.llm.embed_dimensions)
schema.add_field("chunk_type", DataType.VARCHAR, max_length=32)
schema.add_field("table_name", DataType.VARCHAR, max_length=128)
schema.add_field("field_name", DataType.VARCHAR, max_length=128)
schema.add_field("text", DataType.VARCHAR, max_length=65535)
schema.add_field("is_valid", DataType.BOOL)
schema.add_field("is_deprecated", DataType.BOOL)
schema.add_field("content_hash", DataType.VARCHAR, max_length=128)
schema.add_field("created_at", DataType.INT64)
# Few-shot 字段允许元数据 chunk 缺省,避免影响现有表/字段向量写入。
schema.add_field("query", DataType.VARCHAR, max_length=4096, nullable=True)
schema.add_field("correct_sql", DataType.VARCHAR, max_length=16384, nullable=True)
schema.add_field("explanation", DataType.VARCHAR, max_length=65535, nullable=True)
schema.add_field("case_id", DataType.VARCHAR, max_length=128, nullable=True)
return schema
def build_nl2sql_index_params():
params = AsyncMilvusClient.prepare_index_params()
params.add_index(
field_name="vector",
index_type="HNSW",
metric_type="COSINE",
params={"M": 16, "efConstruction": 200},
)
return params
async def ensure_nl2sql_database(milvus_client: AsyncMilvusClient | None = None) -> None:
if milvus_client is None:
from config.database.milvus import client as configured_milvus_client
milvus_client = configured_milvus_client()
client = milvus_client
databases = await client.list_databases()
if NL2SQL_DATABASE_NAME not in databases:
await client.create_database(NL2SQL_DATABASE_NAME)
async def ensure_nl2sql_collection(milvus_client: AsyncMilvusClient | None = None) -> None:
if milvus_client is None:
from config.database.milvus import client as configured_milvus_client
milvus_client = configured_milvus_client()
client = milvus_client
await ensure_nl2sql_database(client)
if await client.has_collection(NL2SQL_COLLECTION):
return
await client.create_collection(
collection_name=NL2SQL_COLLECTION,
schema=build_nl2sql_schema(),
index_params=build_nl2sql_index_params(),
)
async def recreate_nl2sql_collection(
milvus_client: AsyncMilvusClient | None = None,
) -> None:
"""删除并按当前 Embedding 维度重新创建 NL2SQL 集合。"""
if milvus_client is None:
from config.database.milvus import client as configured_milvus_client
milvus_client = configured_milvus_client()
client = milvus_client
await ensure_nl2sql_database(client)
if await client.has_collection(NL2SQL_COLLECTION):
await client.drop_collection(collection_name=NL2SQL_COLLECTION)
await client.create_collection(
collection_name=NL2SQL_COLLECTION,
schema=build_nl2sql_schema(),
index_params=build_nl2sql_index_params(),
)