feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -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(),
|
||||
)
|
||||
Reference in New Issue
Block a user