Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions backend/alembic/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,14 +24,14 @@

# from apps.system.models.user import SQLModel # noqa
# from apps.settings.models.setting_models import SQLModel
#from apps.chat.models.chat_model import SQLModel
from apps.terminology.models.terminology_model import SQLModel
from sqlbot_xpack.custom_prompt.models.custom_prompt_model import SQLModel
from apps.chat.models.chat_model import SQLModel
#from apps.terminology.models.terminology_model import SQLModel
#from sqlbot_xpack.custom_prompt.models.custom_prompt_model import SQLModel
#from apps.data_training.models.data_training_model import SQLModel
# from apps.dashboard.models.dashboard_model import SQLModel
from common.core.config import settings # noqa
#from apps.datasource.models.datasource import SQLModel
from apps.system.models.system_model import SQLModel
#from apps.system.models.system_model import SQLModel

target_metadata = SQLModel.metadata

Expand Down
31 changes: 31 additions & 0 deletions backend/alembic/versions/072_chat_record.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
"""072_chat_record

Revision ID: f898322341a5
Revises: a2e2ecfa5a9c
Create Date: 2026-09-14 15:54:47.921806

"""
from alembic import op
import sqlalchemy as sa
import sqlmodel.sql.sqltypes
from sqlalchemy.dialects import postgresql

# revision identifiers, used by Alembic.
revision = 'f898322341a5'
down_revision = 'a2e2ecfa5a9c'
branch_labels = None
depends_on = None


def upgrade():
# ### commands auto generated by Alembic - please adjust! ###
op.add_column('chat_record', sa.Column('extracted_keywords', sa.Text(), nullable=True, comment='关键词'))
op.add_column('chat_record', sa.Column('expanded_keywords', sa.Text(), nullable=True, comment='扩展后的关键词'))
# ### end Alembic commands ###


def downgrade():
# ### commands auto generated by Alembic - please adjust! ###
op.drop_column('chat_record', 'extracted_keywords')
op.drop_column('chat_record', 'expanded_keywords')
# ### end Alembic commands ###
13 changes: 13 additions & 0 deletions backend/apps/chat/curd/chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -933,6 +933,19 @@ def save_sql_answer(session: SessionDep, record_id: int, answer: str) -> ChatRec
return record


def save_extracted_keywords(session: SessionDep, record_id: int, expanded_keywords: str,
extracted_keywords: str = None) -> None:
"""保存提取的关键词到 record,供后续独立接口复用。"""
if not record_id:
return
values = {'expanded_keywords': expanded_keywords}
if extracted_keywords:
values['extracted_keywords'] = extracted_keywords
stmt = update(ChatRecord).where(and_(ChatRecord.id == record_id)).values(**values)
session.execute(stmt)
session.commit()


def save_analysis_answer(session: SessionDep, record_id: int, answer: str = '') -> ChatRecord:
if not record_id:
raise Exception("Record id cannot be None")
Expand Down
9 changes: 9 additions & 0 deletions backend/apps/chat/models/chat_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@ class OperationEnum(Enum):
FILTER_CUSTOM_PROMPT = '11'
EXECUTE_SQL = '12'
GENERATE_PICTURE = '13'
EXTRACT_KEYWORDS = '14'


class ChatFinishStep(Enum):
Expand Down Expand Up @@ -129,6 +130,8 @@ class ChatRecord(SQLModel, table=True):
analysis_record_id: int = Field(sa_column=Column(BigInteger, nullable=True))
predict_record_id: int = Field(sa_column=Column(BigInteger, nullable=True))
regenerate_record_id: int = Field(sa_column=Column(BigInteger, nullable=True))
extracted_keywords: str = Field(sa_column=Column(Text, nullable=True))
expanded_keywords: str = Field(sa_column=Column(Text, nullable=True))


class ChatRecordResult(BaseModel):
Expand Down Expand Up @@ -238,6 +241,8 @@ class AiModelQuestion(BaseModel):
question: str = None
ai_modal_id: int = None
ai_modal_name: str = None # Specific model name
extracted_keywords: str = "" # LLM 提取的核心业务关键词(逗号分隔,未经扩展)
expanded_keywords: str = "" # 经术语同义词扩展后的关键词(用于表匹配)
engine: str = ""
db_schema: str = ""
sql: str = ""
Expand All @@ -255,6 +260,10 @@ class AiModelQuestion(BaseModel):
sample_data: str = ""
sqlbot_name: str = "SQLBot"

def extract_keywords_sys_prompt(self) -> str:
return get_sql_template().get("extract_keywords").format(lang=self.lang,
sqlbot_name=self.sqlbot_name)

def sql_sys_question(self, db_type: Union[str, DB], enable_query_limit: bool = True):
templates: dict[str, str] = {}
_sql_template = get_sql_example_template(db_type)
Expand Down
141 changes: 137 additions & 4 deletions backend/apps/chat/task/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import concurrent
import json
import os
import re
import traceback
import urllib.parse
import warnings
Expand Down Expand Up @@ -30,7 +31,7 @@
from apps.chat.curd.chat import save_question, save_sql_answer, save_sql, \
save_error_message, save_sql_exec_data, save_chart_answer, save_chart, \
finish_record, save_analysis_answer, save_predict_answer, save_predict_data, \
save_select_datasource_answer, save_recommend_question_answer, \
save_select_datasource_answer, save_recommend_question_answer, save_extracted_keywords, \
get_old_questions, save_analysis_predict_record, rename_chat, get_chart_config, \
get_chat_chart_data, list_generate_sql_logs, list_generate_chart_logs, start_log, end_log, \
get_last_execute_sql_error, format_json_data, format_chart_fields, get_chat_brief_generate, get_chat_predict_data, \
Expand All @@ -48,6 +49,7 @@
from apps.system.crud.parameter_manage import get_groups
from apps.system.crud.user import user_ws_list
from apps.system.schemas.system_schema import AssistantOutDsSchema
from apps.terminology.curd.terminology import expand_with_terminology
from apps.terminology.curd.terminology import get_terminology_template
from common.core.config import settings
from common.core.db import engine
Expand Down Expand Up @@ -378,14 +380,16 @@ def filter_terminology_template(self, _session: Session, oid: int = None, ds_id:
calculate_oid = self.current_assistant.oid if self.current_assistant.type != 4 else self.oid
if self.current_assistant.type == 1:
calculate_ds_id = None
# 使用原始提取的关键词(未扩展)进行术语匹配,避免同义词重复查询
match_text = self.chat_question.extracted_keywords or self.chat_question.question
if self.current_assistant and self.current_assistant.type == 1:
self.chat_question.terminologies, term_list = get_terminology_template(_session,
self.chat_question.question,
match_text,
calculate_oid,
None, self.current_assistant.id)
else:
self.chat_question.terminologies, term_list = get_terminology_template(_session,
self.chat_question.question,
match_text,
calculate_oid,
calculate_ds_id)

Expand Down Expand Up @@ -447,17 +451,108 @@ def filter_training_template(self, _session: Session, oid: int = None, ds_id: in
OperationEnum.FILTER_SQL_EXAMPLE],
full_message=example_list)

def extract_keywords(self, _session: Session) -> tuple[bool, str]:
"""从用户问题中提取核心业务实体词。

使用 LLM 提取名词性业务实体,忽略时间词、动作词、修饰词。
支持上下文历史,可根据对话上下文理解用户意图。
提取结果可复用于表匹配、术语扩展等环节。

Returns:
(is_error, result) 元组:
- (False, keywords) 成功,关键词逗号分隔
- (True, error_message) 错误(LLM 失败或非查数据意图)
"""
system_prompt = self.chat_question.extract_keywords_sys_prompt()

keywords_msg: List[Union[BaseMessage, dict[str, Any]]] = []
keywords_msg.append(SystemPromptMessage(content=system_prompt))

# 加载上下文历史(仅取用户的历史提问,不含 AI 生成的 SQL 等回复)
last_sql_messages: List[dict[str, Any]] = self.generate_sql_logs[-1].messages if len(
self.generate_sql_logs) > 0 else []
if self.chat_question.regenerate_record_id:
_temp_log = next(
filter(lambda obj: obj.pid == self.chat_question.regenerate_record_id, self.generate_sql_logs), None)
last_sql_messages: List[dict[str, Any]] = _temp_log.messages if _temp_log else []

# 排除系统提示词和 AI 回复,只保留用户的历史提问
last_user_messages = [
obj for obj in last_sql_messages
if obj.get("sqlbot_system") != True and obj.get('type') == 'human'
]

if last_user_messages:
last_rounds = get_last_conversation_rounds(last_user_messages, rounds=self.base_message_round_count_limit)
for _msg_dict in last_rounds:
content = _msg_dict.get('content', '')
# 提取 <user-question> 标签内的用户原始提问,排除 error-msg 等干扰信息
match = re.search(r'<user-question>(.*?)</user-question>', content, re.DOTALL)
if match:
question_text = match.group(1).strip()
if question_text:
keywords_msg.append(HumanMessage(content=question_text))

# 当前问题
keywords_msg.append(HumanMessage(content=self.chat_question.question))

self.current_logs[OperationEnum.EXTRACT_KEYWORDS] = start_log(
session=_session,
ai_modal_id=self.chat_question.ai_modal_id,
ai_modal_name=self.chat_question.ai_modal_name,
operate=OperationEnum.EXTRACT_KEYWORDS,
record_id=self.record.id,
full_message=[{'type': msg.type, 'content': msg.content} for msg in keywords_msg]
)

full_thinking_text = ''
full_text = ''
token_usage = {}
res = process_stream(self.llm.stream(keywords_msg), token_usage)
for chunk in res:
if chunk.get('content'):
full_text += chunk.get('content')
if chunk.get('reasoning_content'):
full_thinking_text += chunk.get('reasoning_content')

result_text = full_text.strip()

keywords_msg.append(AIMessage(result_text))
self.current_logs[OperationEnum.EXTRACT_KEYWORDS] = end_log(
session=_session,
log=self.current_logs[OperationEnum.EXTRACT_KEYWORDS],
full_message=[{'type': msg.type, 'content': msg.content} for msg in keywords_msg],
reasoning_content=full_thinking_text if full_thinking_text else None,
token_usage=token_usage
)

if result_text.startswith('ERROR:NOT_DATA_QUERY'):
# return True, '您的问题似乎不是查询数据的问题,请提出与数据查询相关的问题。'
return False, ''

if result_text == 'EMPTY' or not result_text:
self.chat_question.extracted_keywords = self.chat_question.question
return False, ''

self.chat_question.extracted_keywords = result_text
return False, result_text

def choose_table_schema(self, _session: Session):
self.current_logs[OperationEnum.CHOOSE_TABLE] = start_log(session=_session,
operate=OperationEnum.CHOOSE_TABLE,
record_id=self.record.id,
local_operation=True)

# 使用扩展后的关键词进行表匹配(包含同义词,匹配更全面)
keywords = self.chat_question.expanded_keywords or self.chat_question.question

self.chat_question.db_schema, tables = self.out_ds_instance.get_db_schema(
self.ds.id, self.chat_question.question) if self.out_ds_instance else get_table_schema(
session=_session,
current_user=self.current_user,
ds=self.ds,
question=self.chat_question.question)
question=self.chat_question.question,
keywords=keywords)

# Get sample data for all tables
if not self.out_ds_instance:
Expand All @@ -479,6 +574,12 @@ def generate_analysis(self, _session: Session):
self.chat_question.data = orjson.dumps(data.get('data')).decode()
analysis_msg: List[Union[BaseMessage, dict[str, Any]]] = []

# 从 record 加载已提取的关键词(SQL 生成阶段已持久化)
if self.record.extracted_keywords:
self.chat_question.extracted_keywords = self.record.extracted_keywords
if self.record.expanded_keywords:
self.chat_question.expanded_keywords = self.record.expanded_keywords

ds_id = self.ds.id if isinstance(self.ds, CoreDatasource) else None

self.filter_terminology_template(_session, self.oid, ds_id)
Expand Down Expand Up @@ -781,6 +882,22 @@ def select_datasource(self, _session: Session):
oid = self.ds.oid if isinstance(self.ds, CoreDatasource) else 1
ds_id = self.ds.id if isinstance(self.ds, CoreDatasource) else None

# 提取关键词(前置步骤,结果复用于术语匹配和表匹配)
if settings.TABLE_EMBEDDING_KEYWORD_ENABLED:
is_error, kw_result = self.extract_keywords(_session)
if is_error:
raise SingleMessageError(kw_result)
if kw_result:
self.chat_question.extracted_keywords = kw_result
SQLBotLogUtil.info(f"提取的关键词: {kw_result}")
# 术语同义词扩展(用于表匹配,不用于术语模板查询)
expanded = expand_with_terminology(kw_result, _session, oid)
self.chat_question.expanded_keywords = expanded
SQLBotLogUtil.info(f"扩展后的关键词: {expanded}")
# 持久化到 record,供后续独立接口(如 generate_analysis)使用
save_extracted_keywords(_session, self.record.id,
expanded_keywords=expanded, extracted_keywords=kw_result)

self.filter_terminology_template(_session, oid, ds_id)

self.filter_training_template(_session, oid, ds_id)
Expand Down Expand Up @@ -1247,6 +1364,22 @@ def run_task(self, in_chat: bool = True, stream: bool = True,
oid = self.ds.oid if isinstance(self.ds, CoreDatasource) else 1
ds_id = self.ds.id if isinstance(self.ds, CoreDatasource) else None

# 提取关键词(前置步骤,结果复用于术语匹配和表匹配)
if settings.TABLE_EMBEDDING_KEYWORD_ENABLED:
is_error, kw_result = self.extract_keywords(_session)
if is_error:
raise SingleMessageError(kw_result)
if kw_result:
self.chat_question.extracted_keywords = kw_result
SQLBotLogUtil.info(f"提取的关键词: {kw_result}")
# 术语同义词扩展(用于表匹配,不用于术语模板查询)
expanded = expand_with_terminology(kw_result, _session, oid)
self.chat_question.expanded_keywords = expanded
SQLBotLogUtil.info(f"扩展后的关键词: {expanded}")
# 持久化到 record,供后续独立接口(如 generate_analysis)使用
save_extracted_keywords(_session, self.record.id,
expanded_keywords=expanded, extracted_keywords=kw_result)

self.filter_terminology_template(_session, oid, ds_id)

self.filter_training_template(_session, oid, ds_id)
Expand Down
22 changes: 18 additions & 4 deletions backend/apps/datasource/crud/datasource.py
Original file line number Diff line number Diff line change
Expand Up @@ -517,7 +517,8 @@ def get_tables_sample_data(session: SessionDep, current_user: CurrentUser, ds: C


def get_table_schema(session: SessionDep, current_user: CurrentUser, ds: CoreDatasource, question: str,
embedding: bool = True, table_list: list[str] = None) -> tuple[str, list]:
embedding: bool = True, table_list: list[str] = None,
keywords: str = None) -> tuple[str, list]:
schema_str = ""
table_objs = get_table_obj_by_ds(session=session, current_user=current_user, ds=ds)
if len(table_objs) == 0:
Expand All @@ -538,17 +539,25 @@ def get_table_schema(session: SessionDep, current_user: CurrentUser, ds: CoreDat
table_comment = ''
if obj.table.custom_comment:
table_comment = obj.table.custom_comment.strip()
if not table_comment and obj.table.table_comment:
table_comment = obj.table.table_comment.strip()
if table_comment == '':
schema_table += '\n[\n'
else:
schema_table += f", {table_comment}\n[\n"

fields_info = []
if obj.fields:
field_list = []
for field in obj.fields:
field_comment = ''
if field.custom_comment:
field_comment = field.custom_comment.strip()
fields_info.append({
"field_name": field.field_name,
"field_comment": (field.field_comment or '').strip(),
"custom_comment": (field.custom_comment or '').strip()
})
if field_comment == '':
field_list.append(f"({field.field_name}:{field.field_type})")
else:
Expand All @@ -557,17 +566,22 @@ def get_table_schema(session: SessionDep, current_user: CurrentUser, ds: CoreDat
schema_table += '\n]\n'

t_obj = {"id": obj.table.id, "table_name": obj.table.table_name, "schema_table": schema_table,
"embedding": obj.table.embedding}
"embedding": obj.table.embedding, "table_comment": table_comment, "fields": fields_info}
tables.append(t_obj)
all_tables.append(t_obj)

# 如果没有符合过滤条件的表,直接返回
if not tables:
return schema_str, []

# do table embedding
# 执行表 embedding 匹配
if embedding and tables and settings.TABLE_EMBEDDING_ENABLED:
tables = calc_table_embedding(tables, question)
try:
tables = calc_table_embedding(tables, question, session=session, oid=current_user.oid,
keywords=keywords)
except ValueError as e:
# LLM 判断意图不是查数据或关键词提取失败
return str(e), []
# splice schema
if tables:
for s in tables:
Expand Down
Loading
Loading