Spaces:
Runtime error
Runtime error
| """Populate plan_day_content rows by generating study aids for personalized plans.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import logging | |
| import os | |
| from typing import Any, Dict, List | |
| import psycopg | |
| from generate_daily_content import ensure_llm, generate_day_payload | |
| LOGGER = logging.getLogger(__name__) | |
| def _resolve_template_plan_id(cur: psycopg.Cursor, plan_id: str) -> str: | |
| cur.execute( | |
| """ | |
| SELECT COALESCE(template_parent_id, id) | |
| FROM study_plans | |
| WHERE id = %s | |
| """, | |
| (plan_id,), | |
| ) | |
| row = cur.fetchone() | |
| if not row: | |
| raise ValueError(f"Plan {plan_id} not found") | |
| return str(row[0]) | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description="Generate flashcards/practice for a study plan") | |
| parser.add_argument("--plan-id", required=True, help="UUID of the plan to enrich") | |
| parser.add_argument("--start-day", type=int, default=1, help="First day number to include") | |
| parser.add_argument("--end-day", type=int, default=0, help="Last day number (0 = end)") | |
| parser.add_argument( | |
| "--database-url", | |
| default=os.getenv("DATABASE_URL", "postgresql://postgres:postgres@localhost:5432/learning"), | |
| help="PostgreSQL connection string", | |
| ) | |
| parser.add_argument( | |
| "--llm-backend", | |
| choices=["openai", "ollama"], | |
| default="openai", | |
| help="LLM provider", | |
| ) | |
| parser.add_argument("--llm-model", default="gpt-4o-mini", help="Model name when backend=openai") | |
| parser.add_argument("--ollama-model", default="llama3.1", help="Model name when backend=ollama") | |
| parser.add_argument("--dry-run", action="store_true", help="Print prompts without writing to DB") | |
| parser.add_argument( | |
| "--overwrite", | |
| action="store_true", | |
| help="Regenerate days even if content already exists", | |
| ) | |
| return parser.parse_args() | |
| def configure_logging() -> None: | |
| logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s") | |
| def fetch_plan_metadata(cur: psycopg.Cursor, plan_id: str) -> Dict[str, Any]: | |
| cur.execute( | |
| """ | |
| SELECT sp.id, b.title | |
| FROM study_plans sp | |
| JOIN books b ON sp.book_id = b.id | |
| WHERE sp.id = %s | |
| """, | |
| (plan_id,), | |
| ) | |
| row = cur.fetchone() | |
| if not row: | |
| raise ValueError(f"Plan {plan_id} not found") | |
| return {"plan_id": row[0], "book_title": row[1]} | |
| def fetch_plan_days( | |
| cur: psycopg.Cursor, | |
| plan_id: str, | |
| start_day: int, | |
| end_day: int, | |
| overwrite: bool, | |
| ) -> List[Dict[str, Any]]: | |
| clauses = ["pd.plan_id = %s"] | |
| params: List[Any] = [plan_id] | |
| if start_day: | |
| clauses.append("pd.day_number >= %s") | |
| params.append(start_day) | |
| if end_day: | |
| clauses.append("pd.day_number <= %s") | |
| params.append(end_day) | |
| query = f""" | |
| SELECT pd.id, pd.day_number, pd.payload, pdc.content | |
| FROM plan_days pd | |
| LEFT JOIN plan_day_content pdc ON pd.id = pdc.plan_day_id | |
| WHERE {' AND '.join(clauses)} | |
| ORDER BY pd.day_number | |
| """ | |
| cur.execute(query, params) # type: ignore | |
| rows = cur.fetchall() | |
| results: List[Dict[str, Any]] = [] | |
| for row in rows: | |
| if not overwrite and row[3] is not None: | |
| continue | |
| results.append( | |
| { | |
| "plan_day_id": row[0], | |
| "day_number": row[1], | |
| "payload": row[2], | |
| } | |
| ) | |
| return results | |
| def upsert_plan_day_content(cur: psycopg.Cursor, plan_day_id: int, content: Dict[str, Any]) -> None: | |
| cur.execute( | |
| """ | |
| INSERT INTO plan_day_content (plan_day_id, content) | |
| VALUES (%s, %s::jsonb) | |
| ON CONFLICT (plan_day_id) | |
| DO UPDATE SET content = EXCLUDED.content; | |
| """, | |
| (plan_day_id, json.dumps(content, ensure_ascii=False)), | |
| ) | |
| def run_plan_enrichment( | |
| plan_id: str, | |
| database_url: str, | |
| start_day: int = 1, | |
| end_day: int = 0, | |
| overwrite: bool = False, | |
| llm_backend: str = "openai", | |
| llm_model: str = "gpt-4o-mini", | |
| ollama_model: str = "llama3.1", | |
| dry_run: bool = False, | |
| ) -> None: | |
| import argparse | |
| cfg = argparse.Namespace(llm_backend=llm_backend, llm_model=llm_model, ollama_model=ollama_model) | |
| llm = None if dry_run else ensure_llm(cfg) | |
| with psycopg.connect(database_url) as conn: | |
| with conn.cursor() as cur: | |
| resolved_plan_id = _resolve_template_plan_id(cur, plan_id) | |
| if resolved_plan_id != plan_id: | |
| LOGGER.info("Resolved plan %s to template %s for enrichment", plan_id, resolved_plan_id) | |
| plan_id = resolved_plan_id | |
| plan_meta = fetch_plan_metadata(cur, plan_id) | |
| plan_days = fetch_plan_days(cur, plan_id, start_day, end_day, overwrite) | |
| if not plan_days: | |
| LOGGER.info("No plan days matched the criteria.") | |
| return | |
| LOGGER.info("Generating content for %s days", len(plan_days)) | |
| for entry in plan_days: | |
| plan_day_id = entry["plan_day_id"] | |
| day_payload = entry["payload"] or {} | |
| if not day_payload: | |
| LOGGER.warning("Skipping plan_day_id=%s due to missing payload", plan_day_id) | |
| continue | |
| content = generate_day_payload(llm, plan_meta["book_title"], day_payload, dry_run) | |
| if dry_run: | |
| continue | |
| upsert_plan_day_content(cur, plan_day_id, content) | |
| LOGGER.info("Saved content for day %s (plan_day_id=%s)", entry["day_number"], plan_day_id) | |
| conn.commit() | |
| conn.commit() | |
| LOGGER.info("Completed enrichment for plan %s", plan_id) | |
| def main() -> None: | |
| configure_logging() | |
| args = parse_args() | |
| run_plan_enrichment( | |
| plan_id=args.plan_id, | |
| database_url=args.database_url, | |
| start_day=args.start_day, | |
| end_day=args.end_day, | |
| overwrite=args.overwrite, | |
| llm_backend=args.llm_backend, | |
| llm_model=args.llm_model, | |
| ollama_model=args.ollama_model, | |
| dry_run=args.dry_run, | |
| ) | |
| if __name__ == "__main__": | |
| main() | |