feat:修改nl2sql功能
This commit is contained in:
+18
-4
@@ -107,18 +107,26 @@ def resolve_semantics(question: str, *, catalog: dict[str, Any] | None = None) -
|
||||
{
|
||||
"term": item["term"],
|
||||
"field": item["fields"][0],
|
||||
"hint": item.get("metric_hint") or item.get("value_hint", ""),
|
||||
"hint": item["metric_hint"],
|
||||
}
|
||||
# metric_hint(指标口径)与 value_hint(字段取值口径)都要进入
|
||||
# SQL 生成上下文;只过滤 metric_hint 会让枚举值提示永远丢失。
|
||||
for item in matched
|
||||
if "metric_hint" in item or "value_hint" in item
|
||||
if item.get("metric_hint")
|
||||
]
|
||||
value_hints = [
|
||||
{
|
||||
"term": item["term"],
|
||||
"field": item["fields"][0],
|
||||
"hint": item["value_hint"],
|
||||
}
|
||||
for item in matched
|
||||
if item.get("value_hint")
|
||||
]
|
||||
return {
|
||||
"terms": [item["term"] for item in matched],
|
||||
"tables": tables,
|
||||
"fields": fields,
|
||||
"metrics": metrics,
|
||||
"value_hints": value_hints,
|
||||
"relationships": list(catalog.get("relationships", [])),
|
||||
}
|
||||
|
||||
@@ -149,6 +157,11 @@ def build_semantic_context(question: str, schema: dict[str, Any]) -> dict[str, A
|
||||
for metric in resolved["metrics"]
|
||||
if any(metric["field"] == field for _, field in schema_fields if _ in tables)
|
||||
]
|
||||
value_hints = [
|
||||
hint
|
||||
for hint in resolved["value_hints"]
|
||||
if any(hint["field"] == field for _, field in schema_fields if _ in tables)
|
||||
]
|
||||
relationships = [
|
||||
relation
|
||||
for relation in resolved["relationships"]
|
||||
@@ -160,5 +173,6 @@ def build_semantic_context(question: str, schema: dict[str, Any]) -> dict[str, A
|
||||
"tables": tables,
|
||||
"fields": fields,
|
||||
"metrics": metrics,
|
||||
"value_hints": value_hints,
|
||||
"relationships": relationships,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user