Skip to content

Commit d17c433

Browse files
authored
Merge pull request #244 from thrcle/feat/kind-aware-prompt
feat: system prompt kind별 그룹화 + disambiguation 개선
2 parents 0f41059 + 85342cf commit d17c433

3 files changed

Lines changed: 173 additions & 42 deletions

File tree

bench/ecommerce_demo.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -31,9 +31,8 @@
3131
from lang2sql.tools.semantic_federation import (
3232
FedEntry,
3333
_kv_key,
34-
_render_effective,
3534
_load_all,
36-
_resolve_term,
35+
_render_effective,
3736
)
3837

3938
# Stable IDs for the demo guild and its two channels.

src/lang2sql/tools/semantic_federation.py

Lines changed: 68 additions & 40 deletions
Original file line numberDiff line numberDiff line change
@@ -400,35 +400,49 @@ def _load_all(store: Any, scope: str) -> dict[str, list[FedEntry]]:
400400
return by_term
401401

402402

403-
def build_prompt_section(store: Any, scope: str, channel_id: str, user_id: str) -> str:
404-
"""현재 채널 기준 narrow→wide lookup 용어 섹션 + 모호 용어 지침 반환."""
405-
by_term = _load_all(store, scope)
406-
407-
if not by_term:
408-
return _AMBIGUOUS_TERM_POLICY
409-
410-
lines: list[str] = []
411-
for term_lower in sorted(by_term):
412-
line = _resolve_term(by_term[term_lower], channel_id, user_id)
413-
if line:
414-
lines.append(line)
415-
416-
header = "## Business Terminology\n" "(lookup 우선순위: 개인 > 채널(팀) > 전사)\n"
417-
body = "\n".join(lines) if lines else "(없음)"
418-
return header + body + "\n\n" + _AMBIGUOUS_TERM_POLICY
419-
403+
_KIND_SQL_HINT: dict[str, str] = {
404+
"metric": "집계 지표 — SELECT/HAVING 절의 집계식으로 사용",
405+
"rule": "비즈니스 규칙 — WHERE 절에 AND 조건으로 추가",
406+
"dimension": "분류 기준 — GROUP BY 또는 SELECT 컬럼으로 활용",
407+
"table": "테이블/엔티티 — FROM/JOIN 대상 선택 시 참고",
408+
}
409+
_KIND_ORDER = ["metric", "dimension", "rule", "table"]
420410

421411
_AMBIGUOUS_TERM_POLICY = """\
422412
## Ambiguous Term Policy
423-
사전에 없는 주관적/모호한 표현(예: 활성화고객, 신규고객, 우량고객)을 발견하면:
424-
1. 현재 DB 스키마 컨텍스트에서 가장 합리적인 해석으로 SQL을 작성하고 실행한다.
425-
2. 쿼리 후 사용한 해석을 명시하고, term_custom 등록 여부와 범위(guild/channel/member)를 사용자에게 묻는다.
426-
예: "'신규고객'을 'users.created_at >= NOW()-30일'로 해석했습니다. 이 정의를 어느 범위로 등록할까요?"
427-
3. 사용자가 범위를 지정하면 term_custom 툴로 즉시 등록한다 (inferred=true).
413+
사전에 없는 주관적/모호한 표현을 발견하면:
414+
1. DB 스키마 기준으로 가장 합리적인 해석으로 SQL을 실행한다.
415+
2. 실행 후 사용한 해석을 명시하고, kind(metric/rule/dimension/table)와 범위(guild/channel/member)를 사용자에게 묻는다.
416+
예: "'신규고객'을 'users.created_at >= NOW()-30일'로 해석했습니다. metric/rule/dimension/table 중 어느 종류이며, 어느 범위로 등록할까요?"
417+
3. 사용자가 지정하면 term_custom 툴로 즉시 등록한다 (inferred=true).
428418
4. inferred=true 엔트리가 이미 있으면 해당 정의를 우선 사용하되, 사용자에게 확정 여부를 확인한다.\
429419
"""
430420

431421

422+
def _resolve_entry(
423+
entries: list[FedEntry], channel_id: str, user_id: str
424+
) -> FedEntry | None:
425+
"""narrow→wide lookup: member > channel > guild. 승리 FedEntry 반환."""
426+
for e in entries:
427+
if e.layer == "member" and e.entity == user_id:
428+
return e
429+
for e in entries:
430+
if e.layer == "channel" and e.entity == channel_id:
431+
return e
432+
for e in entries:
433+
if e.layer == "guild":
434+
return e
435+
return None
436+
437+
438+
def _tag_for(e: FedEntry) -> str:
439+
if e.layer == "member":
440+
return f"개인:{e.entity}"
441+
if e.layer == "channel":
442+
return "채널"
443+
return "전사"
444+
445+
432446
def _fmt_entry(e: FedEntry, tag: str) -> str:
433447
syns = ", ".join(e.synonyms)
434448
syn_str = f" (= {syns})" if syns else ""
@@ -439,24 +453,38 @@ def _fmt_entry(e: FedEntry, tag: str) -> str:
439453
)
440454

441455

442-
def _resolve_term(entries: list[FedEntry], channel_id: str, user_id: str) -> str:
443-
"""narrow→wide lookup: member > channel > guild."""
444-
# 1. 개인 오버라이드
445-
for e in entries:
446-
if e.layer == "member" and e.entity == user_id:
447-
return _fmt_entry(e, f"개인:{user_id}")
456+
def build_prompt_section(store: Any, scope: str, channel_id: str, user_id: str) -> str:
457+
"""kind별로 그룹화된 시멘틱 용어 섹션 + 모호 용어 지침 반환."""
458+
by_term = _load_all(store, scope)
448459

449-
# 2. 이 채널 정의
450-
for e in entries:
451-
if e.layer == "channel" and e.entity == channel_id:
452-
return _fmt_entry(e, "채널")
460+
if not by_term:
461+
return _AMBIGUOUS_TERM_POLICY
453462

454-
# 3. 전사 공통
455-
for e in entries:
456-
if e.layer == "guild":
457-
return _fmt_entry(e, "전사")
463+
groups: dict[str, list[str]] = {}
464+
for term_lower in sorted(by_term):
465+
e = _resolve_entry(by_term[term_lower], channel_id, user_id)
466+
if e is None:
467+
continue
468+
k = e.kind if e.kind in _KIND_SQL_HINT else ""
469+
groups.setdefault(k, []).append(_fmt_entry(e, _tag_for(e)))
470+
471+
if not groups:
472+
return _AMBIGUOUS_TERM_POLICY
473+
474+
parts: list[str] = [
475+
"## Business Terminology\n(lookup 우선순위: 개인 > 채널(팀) > 전사)\n"
476+
]
477+
for kind in _KIND_ORDER + [""]:
478+
lines = groups.get(kind, [])
479+
if not lines:
480+
continue
481+
if kind:
482+
parts.append(f"### {kind.capitalize()}s — {_KIND_SQL_HINT[kind]}")
483+
else:
484+
parts.append("### 기타")
485+
parts.extend(lines)
458486

459-
return ""
487+
return "\n".join(parts) + "\n\n" + _AMBIGUOUS_TERM_POLICY
460488

461489

462490
def _render_effective(store: Any, scope: str, channel_id: str, user_id: str) -> str:
@@ -467,9 +495,9 @@ def _render_effective(store: Any, scope: str, channel_id: str, user_id: str) ->
467495

468496
lines = ["**Business Terminology — 현재 채널 기준 유효 정의**\n"]
469497
for term_lower in sorted(by_term):
470-
line = _resolve_term(by_term[term_lower], channel_id, user_id)
471-
if line:
472-
lines.append(line)
498+
e = _resolve_entry(by_term[term_lower], channel_id, user_id)
499+
if e:
500+
lines.append(_fmt_entry(e, _tag_for(e)))
473501

474502
if len(lines) == 1:
475503
lines.append("(이 채널에 적용되는 용어 정의가 없습니다)")

tests/test_semantic.py

Lines changed: 104 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -166,3 +166,107 @@ def test_fmt_entry_shows_kind_badge() -> None:
166166
)
167167
rendered = _fmt_entry(entry, "전사")
168168
assert "`metric`" in rendered
169+
170+
171+
# ---------------------------------------------------------------------------
172+
# PR4: kind-grouped prompt section + disambiguation policy
173+
# ---------------------------------------------------------------------------
174+
175+
176+
def _seed(store: SqliteStore, scope: str, entries: list[FedEntry]) -> None:
177+
for e in entries:
178+
store.kv_set(scope, _kv_key(e.term, e.layer, e.entity), e.to_json())
179+
180+
181+
def test_prompt_section_groups_by_kind() -> None:
182+
store = SqliteStore()
183+
_seed(
184+
store,
185+
"g1",
186+
[
187+
FedEntry("월매출", "guild", "", "SUM(orders.amount)", kind="metric"),
188+
FedEntry("환불제외", "guild", "", "status != refunded", kind="rule"),
189+
FedEntry("고객등급", "guild", "", "users.tier", kind="dimension"),
190+
],
191+
)
192+
section = build_prompt_section(store, "g1", "c1", "u1")
193+
194+
assert "### Metrics" in section
195+
assert "### Rules" in section
196+
assert "### Dimensions" in section
197+
assert "월매출" in section
198+
assert "환불제외" in section
199+
assert "고객등급" in section
200+
201+
202+
def test_prompt_section_kind_headers_contain_sql_hint() -> None:
203+
store = SqliteStore()
204+
_seed(
205+
store,
206+
"g1",
207+
[FedEntry("월매출", "guild", "", "SUM(orders.amount)", kind="metric")],
208+
)
209+
section = build_prompt_section(store, "g1", "c1", "u1")
210+
211+
assert "SELECT/HAVING" in section
212+
assert "### Metrics" in section
213+
214+
215+
def test_prompt_section_unknown_kind_goes_to_기타() -> None:
216+
store = SqliteStore()
217+
_seed(store, "g1", [FedEntry("알수없음", "guild", "", "정의 없음", kind="")])
218+
section = build_prompt_section(store, "g1", "c1", "u1")
219+
220+
assert "### 기타" in section
221+
assert "알수없음" in section
222+
223+
224+
def test_prompt_section_skips_empty_kind_groups() -> None:
225+
store = SqliteStore()
226+
_seed(
227+
store,
228+
"g1",
229+
[FedEntry("월매출", "guild", "", "SUM(orders.amount)", kind="metric")],
230+
)
231+
section = build_prompt_section(store, "g1", "c1", "u1")
232+
233+
assert "### Rules" not in section
234+
assert "### Dimensions" not in section
235+
236+
237+
def test_ambiguous_policy_mentions_kind() -> None:
238+
store = SqliteStore()
239+
section = build_prompt_section(store, "g1", "c1", "u1")
240+
assert "metric/rule/dimension/table" in section
241+
242+
243+
def test_resolve_entry_member_wins_over_channel() -> None:
244+
from lang2sql.tools.semantic_federation import _resolve_entry
245+
246+
entries = [
247+
FedEntry("t", "guild", "", "guild-def"),
248+
FedEntry("t", "channel", "c1", "channel-def"),
249+
FedEntry("t", "member", "u1", "member-def"),
250+
]
251+
result = _resolve_entry(entries, "c1", "u1")
252+
assert result is not None
253+
assert result.definition == "member-def"
254+
255+
256+
def test_resolve_entry_channel_wins_over_guild() -> None:
257+
from lang2sql.tools.semantic_federation import _resolve_entry
258+
259+
entries = [
260+
FedEntry("t", "guild", "", "guild-def"),
261+
FedEntry("t", "channel", "c1", "channel-def"),
262+
]
263+
result = _resolve_entry(entries, "c1", "u1")
264+
assert result is not None
265+
assert result.definition == "channel-def"
266+
267+
268+
def test_resolve_entry_returns_none_when_no_match() -> None:
269+
from lang2sql.tools.semantic_federation import _resolve_entry
270+
271+
entries = [FedEntry("t", "channel", "other-channel", "def")]
272+
assert _resolve_entry(entries, "c1", "u1") is None

0 commit comments

Comments
 (0)