| """Session orchestration for the Hugging Face Gradio Space.""" |
|
|
| from __future__ import annotations |
|
|
| import sqlite3 |
| import sys |
| import threading |
| from dataclasses import dataclass |
| from pathlib import Path |
|
|
| APP_ROOT = Path(__file__).resolve().parent |
| if str(APP_ROOT) not in sys.path: |
| sys.path.insert(0, str(APP_ROOT)) |
|
|
| from backend.llm_service import rank_result_to_suggestion |
| from backend.schemas import SessionAnswerIn, SessionAnswersIn, SuggestionOut |
| from backend.services import mark_unavailable, session_to_filters, update_session |
| from backend.session_store import SessionStore, SuggestionSession |
| from game_genie.llm.gemini import ask_about_game, rank_from_filters |
| from game_genie.search import autocomplete_games, count_games, get_game_by_id |
|
|
| from config import ensure_database, get_database_path |
| from llm.hf_client import HFInferenceClient, create_hf_client |
|
|
| |
| GENRE_OPTIONS: dict[str, dict] = { |
| "fun": { |
| "category_ids": [1002, 1030, 1079, 1041], |
| "mechanic_ids": [2004, 2040, 2047, 2035], |
| "max_weight": 2.2, |
| }, |
| "strategy": { |
| "category_ids": [1009, 1015, 1021, 1029], |
| "mechanic_ids": [2080, 2041, 2002, 2012], |
| "min_weight": 2.5, |
| }, |
| "social": { |
| "category_ids": [1030, 1023, 1039, 1037], |
| "mechanic_ids": [2014, 2017, 2020, 2027], |
| "max_weight": 2.5, |
| "max_play_time": 45, |
| }, |
| "adventure": { |
| "category_ids": [1022, 1020, 1010, 1024], |
| "mechanic_ids": [2023, 2011, 2072, 2026], |
| "min_weight": 2.0, |
| "max_weight": 3.5, |
| }, |
| } |
|
|
| WEIGHT_MAP = { |
| "light": (None, 2.0), |
| "medium": (2.0, 3.0), |
| "heavy": (3.0, None), |
| } |
|
|
|
|
| def _open_connection(db_path: Path) -> sqlite3.Connection: |
| connection = sqlite3.connect(db_path, check_same_thread=False) |
| connection.row_factory = sqlite3.Row |
| return connection |
|
|
|
|
| @dataclass |
| class GameChoice: |
| id: int |
| label: str |
|
|
|
|
| class GameGenieEngine: |
| """Thin wrapper around existing backend search and LLM services.""" |
|
|
| def __init__( |
| self, |
| *, |
| database_path: Path | None = None, |
| hf_client: HFInferenceClient | None = None, |
| session_store: SessionStore | None = None, |
| ) -> None: |
| self._db_path = ensure_database(database_path or get_database_path()) |
| self._local = threading.local() |
| self._hf_client = hf_client if hf_client is not None else create_hf_client() |
| self._sessions = session_store or SessionStore() |
|
|
| @property |
| def _connection(self) -> sqlite3.Connection: |
| conn = getattr(self._local, "connection", None) |
| if conn is None: |
| conn = _open_connection(self._db_path) |
| self._local.connection = conn |
| return conn |
|
|
| @property |
| def hf_client(self) -> HFInferenceClient | None: |
| return self._hf_client |
|
|
| def start_session(self) -> SuggestionSession: |
| return self._sessions.create() |
|
|
| def get_session(self, session_id: str) -> SuggestionSession | None: |
| return self._sessions.get(session_id) |
|
|
| def apply_answers( |
| self, |
| session: SuggestionSession, |
| *, |
| players: int | None = None, |
| max_play_time: int | None = None, |
| weight_level: str | None = None, |
| max_player_age: int | None = None, |
| genre_id: str | None = None, |
| free_text: str | None = None, |
| include_expansions: bool = False, |
| reference_game_ids: list[int] | None = None, |
| excluded_game_ids: list[int] | None = None, |
| ) -> None: |
| min_weight, max_weight = WEIGHT_MAP.get(weight_level or "", (None, None)) |
| genre = GENRE_OPTIONS.get(genre_id or "", {}) |
|
|
| if weight_level: |
| resolved_min_weight = min_weight |
| resolved_max_weight = max_weight |
| else: |
| resolved_min_weight = genre.get("min_weight") |
| resolved_max_weight = genre.get("max_weight") |
|
|
| resolved_max_play_time = max_play_time if max_play_time is not None else genre.get("max_play_time") |
|
|
| answers = SessionAnswersIn( |
| players=players, |
| max_play_time=resolved_max_play_time, |
| min_weight=resolved_min_weight, |
| max_weight=resolved_max_weight, |
| max_player_age=max_player_age, |
| category_ids=genre.get("category_ids", []), |
| mechanic_ids=genre.get("mechanic_ids", []), |
| free_text=free_text.strip() if free_text and free_text.strip() else None, |
| include_expansions=include_expansions, |
| ) |
| update_session( |
| session, |
| SessionAnswerIn( |
| answers=answers, |
| reference_game_ids=reference_game_ids, |
| excluded_game_ids=excluded_game_ids, |
| ), |
| ) |
|
|
| def get_suggestions(self, session: SuggestionSession) -> SuggestionOut: |
| filters = session_to_filters(session) |
| total = count_games(self._connection, filters) |
| rank = rank_from_filters( |
| self._connection, |
| filters, |
| reference_game_ids=session.reference_game_ids, |
| free_text=session.answers.free_text, |
| client=self._hf_client, |
| ) |
| message = None |
| if total == 0: |
| message = "No matching games found. Try relaxing your filters." |
| elif total < 3: |
| message = f"Only {total} matching games found." |
| suggestion = rank_result_to_suggestion( |
| self._connection, |
| rank, |
| filters, |
| message=message, |
| ) |
| session.last_suggestion = suggestion |
| return suggestion |
|
|
| def mark_unavailable(self, session: SuggestionSession, game_id: int) -> SuggestionOut: |
| suggestion = mark_unavailable(self._connection, session, game_id) |
| return suggestion |
|
|
| def refine( |
| self, |
| session: SuggestionSession, |
| *, |
| rejected_game_id: int, |
| rejection_chips: list[str] | None = None, |
| rejection_reason: str | None = None, |
| ) -> SuggestionOut: |
| if rejected_game_id not in session.rejected_game_ids: |
| session.rejected_game_ids.append(rejected_game_id) |
|
|
| chips = rejection_chips or [] |
| if rejection_reason: |
| chips = [*chips, rejection_reason] |
| rejection_context = ", ".join(chips) if chips else None |
|
|
| filters = session_to_filters(session) |
| total = count_games(self._connection, filters) |
| rank = rank_from_filters( |
| self._connection, |
| filters, |
| reference_game_ids=session.reference_game_ids, |
| free_text=session.answers.free_text, |
| rejection_context=rejection_context, |
| client=self._hf_client, |
| ) |
| message = None |
| if total == 0: |
| message = "No matching games found. Try relaxing your filters." |
| elif total < 3: |
| message = f"Only {total} matching games found." |
| suggestion = rank_result_to_suggestion( |
| self._connection, |
| rank, |
| filters, |
| message=message, |
| ) |
| session.last_suggestion = suggestion |
| return suggestion |
|
|
| def ask(self, game_id: int, question: str): |
| return ask_about_game( |
| self._connection, |
| game_id, |
| question, |
| client=self._hf_client, |
| ) |
|
|
| def search_games(self, query: str, *, limit: int = 10) -> list[GameChoice]: |
| if len(query.strip()) < 2: |
| return [] |
| results = autocomplete_games(self._connection, query.strip(), limit=limit) |
| return [ |
| GameChoice(id=game.id, label=game.display_name or game.name) |
| for game in results |
| ] |
|
|
| def choices_for_dropdown(self, query: str) -> list[tuple[str, int]]: |
| return [(choice.label, choice.id) for choice in self.search_games(query)] |
|
|
| def labels_for_ids(self, game_ids: list[int]) -> list[str]: |
| labels: list[str] = [] |
| for game_id in game_ids: |
| game = get_game_by_id(self._connection, game_id) |
| if game is not None: |
| labels.append(game.display_name or game.name) |
| return labels |
|
|
|
|
| _default_engine: GameGenieEngine | None = None |
|
|
|
|
| def get_engine() -> GameGenieEngine: |
| global _default_engine |
| if _default_engine is None: |
| _default_engine = GameGenieEngine() |
| return _default_engine |
|
|