Commit 69da1c59 authored by Kantz's avatar Kantz
Browse files

subsection aus Tasks für das Retrieval nutzen

parent ba5fb86b
from __future__ import annotations
from app.LLM_services import task_hint_LLM
from app.deterministic_services import context_store, task_catalog
import app.config as config
from app.deterministic_services import (
context_store,
embedding_provider,
retrieval_store,
task_catalog,
)
from app.deterministic_services.orchestrators import orchestrator_base as base
def _ensure_context_task_fields(state: base.ChatState, query_text: str) -> tuple[str, str] | None:
store_new = context_store.context_store_new
has_task = bool(store_new.get_task(state.sheet))
......@@ -12,8 +19,13 @@ def _ensure_context_task_fields(state: base.ChatState, query_text: str) -> tuple
if has_task and has_hints and has_solution:
selected = task_catalog.get_selected_task_ids(state.sheet)
if selected[0] and selected[1]:
return selected[0], selected[1]
return None
was_selected = task_catalog.select_task_by_ids(
state.sheet,
selected[0],
selected[1],
)
if was_selected:
return selected[0], selected[1]
sources_text = "\n".join([source.to_string() for source in context_store.get_retrieval(state.sheet)])
selection = task_catalog.select_task_for_context(
......@@ -43,13 +55,48 @@ def _ensure_context_task_fields(state: base.ChatState, query_text: str) -> tuple
return selected_file_id, selected_task_id
def _retrieve_context_for_task(state: base.ChatState, query_text: str) -> None:
refs = task_catalog.get_selected_task_subsection_refs(state.sheet)
if not refs:
return
def _retrieve() -> dict:
embedder, cache_hit = embedding_provider.get_embedder()
sources = retrieval_store.retrieve_with_subsections(
pg_url=config.get_postgres_url(),
embedder=embedder,
query=query_text,
subsection_refs=refs,
k=8,
expand_links=True,
)
context_store.update_retrieval_context(state.sheet, sources)
return {
"subsection_refs": refs,
"source_count": len(sources),
"embedder_cache_hit": cache_hit,
}
base.log_timed_call(
state.tool_log,
"retrieve_context_with_task_subsections",
{
"query": query_text,
"subsection_refs": refs,
},
_retrieve,
)
def _on_bootstrap(state: base.ChatState, query_text: str) -> None:
base.bootstrap_retrieval(state.sheet, query_text, state.tool_log)
_ensure_context_task_fields(state, query_text)
_retrieve_context_for_task(state, query_text)
def _on_turn_logic(state: base.ChatState) -> None:
_ensure_context_task_fields(state, state.last_user)
_retrieve_context_for_task(state, state.last_user)
def _on_build_reply(state: base.ChatState) -> str | None:
......
......@@ -8,6 +8,7 @@ from typing import Any
from app.deterministic_services import context_store
TASKS_DIR = Path(__file__).resolve().parents[2] / "sources" / "tasks"
SUBSECTION_MAP_PATH = TASKS_DIR / "_subsection_map.json"
def _normalize_text(value: str) -> str:
......@@ -18,6 +19,58 @@ def _tokenize(value: str) -> set[str]:
return set(re.findall(r"[a-z0-9_]+", _normalize_text(value)))
def _normalize_subsection_key(value: str) -> str:
collapsed = re.sub(r"[-_]+", " ", value.strip().lower())
return re.sub(r"\s+", " ", collapsed)
def _parse_subsection_ref(value: str) -> tuple[int, int] | None:
match = re.match(r"^\s*(\d+)\s*[:.]\s*(\d+)\s*$", str(value))
if not match:
return None
return int(match.group(1)), int(match.group(2))
def load_subsection_map(path: Path = SUBSECTION_MAP_PATH) -> dict[str, tuple[int, int]]:
if not path.exists():
return {}
try:
content = json.loads(path.read_text(encoding="utf-8"))
except Exception:
return {}
if not isinstance(content, dict):
return {}
mapped: dict[str, tuple[int, int]] = {}
for raw_key, raw_ref in content.items():
key = _normalize_subsection_key(str(raw_key))
parsed = _parse_subsection_ref(str(raw_ref))
if not key or parsed is None:
continue
mapped[key] = parsed
return mapped
def _resolve_task_subsection_refs(
task_file: dict[str, Any],
subsection_map: dict[str, tuple[int, int]] | None = None,
) -> list[tuple[int, int]]:
mapping = subsection_map if subsection_map is not None else load_subsection_map()
subsections = task_file.get("subsections", [])
if not isinstance(subsections, list):
return []
refs: set[tuple[int, int]] = set()
for subsection in subsections:
key = _normalize_subsection_key(str(subsection))
if not key:
continue
ref = mapping.get(key)
if ref is not None:
refs.add((int(ref[0]), int(ref[1])))
return sorted(refs)
def _match_score(query_text: str, candidate_text: str) -> int:
query_tokens = _tokenize(query_text)
if not query_tokens:
......@@ -87,6 +140,8 @@ def set_selected_task(
store_new.set_solution(sheet, solution)
sheet["task_file_id"] = str(task_file.get("_file_id", ""))
sheet["task_id"] = str(task_entry.get("id", "")).zfill(2)
refs = _resolve_task_subsection_refs(task_file)
sheet["task_subsection_refs"] = [[sec, sub] for sec, sub in refs]
def select_task_by_ids(
......@@ -112,6 +167,21 @@ def get_selected_task_ids(sheet: dict[str, Any]) -> tuple[str | None, str | None
return (file_id or None, task_id or None)
def get_selected_task_subsection_refs(sheet: dict[str, Any]) -> list[tuple[int, int]]:
refs_raw = sheet.get("task_subsection_refs", [])
if not isinstance(refs_raw, list):
return []
refs: set[tuple[int, int]] = set()
for item in refs_raw:
if isinstance(item, (list, tuple)) and len(item) >= 2:
try:
refs.add((int(item[0]), int(item[1])))
except Exception:
continue
return sorted(refs)
def select_task_for_context(
sheet: dict[str, Any],
query_text: str,
......
Supports Markdown
0% or .
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment