Team Ai
Apppublic

Hariprita/nl2sql-openenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
tasks.py270 linesDownload Raw Back to server
1from typing import List, Dict, Tuple, Optional2 3# ---------------------------------------------------------------------------4# Task definitions5# ---------------------------------------------------------------------------6TASKS = [7    {8        "id": "simple_select",9        "difficulty": "easy",10        "question": (11            "List the full name and city of every customer from the United States. "12            "Order the results alphabetically by name."13        ),14        "hint": "Look at the customers table. Filter by country = 'United States'.",15    },16    {17        "id": "join_aggregation",18        "difficulty": "medium",19        "question": (20            "What is the total revenue generated by each product category? "21            "Show the category name and total revenue, ordered from highest to lowest revenue."22        ),23        "hint": (24            "Join order_items with products on product_id. "25            "Revenue = SUM(quantity * unit_price). GROUP BY category."26        ),27    },28    {29        "id": "window_ranking",30        "difficulty": "hard",31        "question": (32            "For each customer who has placed at least one order, show their name, "33            "their most recent order date, and their rank by total spending "34            "(rank 1 = highest total spending). Order the results by rank."35        ),36        "hint": (37            "Use RANK() OVER (ORDER BY total_spent DESC) — supported in SQLite 3.25+. "38            "Or use a correlated subquery: "39            "SELECT COUNT(*) + 1 FROM ... WHERE total > current_total."40        ),41    },42]43 44 45def get_task_by_id(task_id: str) -> dict:46    for t in TASKS:47        if t["id"] == task_id:48            return t49    raise ValueError(f"Unknown task: {task_id!r}")50 51 52# ---------------------------------------------------------------------------53# Grader: simple_select54# ---------------------------------------------------------------------------55_US_NAMES = {56    "alice johnson", "bob smith", "carol white",57    "grace lee", "ivy chen", "jack taylor",58}59_US_CITIES = {60    "new york", "los angeles", "chicago",61    "houston", "san francisco", "miami",62}63 64 65def grade_simple_select(66    rows: List[dict],67    columns: List[str],68    error: Optional[str],69) -> Tuple[float, str]:70    """71    Expected: 6 US customers — name + city — ordered alphabetically by name.72    """73    if error:74        return 0.0, f"Query error: {error}"75    if not rows:76        return 0.0, "No rows returned. Check your WHERE clause and table name."77 78    returned_names = {str(r.get("name", "")).lower() for r in rows}79    returned_cities = {str(r.get("city", "")).lower() for r in rows}80 81    # Try alternate column names82    if not returned_names - {""} :83        for col in columns:84            if "name" in col.lower():85                returned_names = {str(r.get(col, "")).lower() for r in rows}86                break87    if not returned_cities - {""}:88        for col in columns:89            if "city" in col.lower():90                returned_cities = {str(r.get(col, "")).lower() for r in rows}91                break92 93    name_overlap = len(returned_names & _US_NAMES) / len(_US_NAMES)94    city_overlap = len(returned_cities & _US_CITIES) / len(_US_CITIES)95 96    # Bonus for correct alphabetical ordering97    name_col = next((c for c in columns if "name" in c.lower()), None)98    ordered = False99    if name_col and len(rows) > 1:100        names = [str(r.get(name_col, "")) for r in rows]101        ordered = names == sorted(names)102 103    base_score = name_overlap * 0.5 + city_overlap * 0.4104    if ordered and base_score > 0.5:105        base_score = min(1.0, base_score + 0.1)106 107    if base_score >= 0.95:108        return 1.0, (109            "Perfect! All 6 US customers returned with correct cities, ordered alphabetically."110        )111    elif base_score >= 0.7:112        return round(base_score, 2), (113            f"Mostly correct. Name coverage: {name_overlap:.0%}, "114            f"city coverage: {city_overlap:.0%}."115        )116    elif base_score >= 0.3:117        return round(base_score, 2), (118            "Partial result. Check your WHERE clause and column selection."119        )120    else:121        return 0.1, "Very few correct results. Re-read the question and schema carefully."122 123 124# ---------------------------------------------------------------------------125# Grader: join_aggregation126# ---------------------------------------------------------------------------127_EXPECTED_CATEGORIES = {"electronics", "clothing", "books"}128 129 130def grade_join_aggregation(131    rows: List[dict],132    columns: List[str],133    error: Optional[str],134) -> Tuple[float, str]:135    """136    Expected: 3 rows — Electronics (highest), Clothing, Books (lowest).137    Agent must JOIN order_items with products and GROUP BY category.138    """139    if error:140        return 0.0, f"Query error: {error}"141    if not rows:142        return 0.0, "No rows returned. You need to JOIN order_items with products."143 144    # Find category and revenue columns145    cat_col = next(146        (c for c in columns if "cat" in c.lower()),147        None,148    )149    rev_col = next(150        (c for c in columns if any(x in c.lower() for x in ["rev", "total", "sum", "amount", "sales"])),151        None,152    )153 154    if not cat_col or not rev_col:155        return 0.2, (156            "Cannot identify category or revenue columns. "157            "Alias them as 'category' and 'total_revenue' in your SELECT."158        )159 160    returned_cats = {str(r.get(cat_col, "")).lower() for r in rows}161    cat_coverage = len(returned_cats & _EXPECTED_CATEGORIES) / len(_EXPECTED_CATEGORIES)162 163    # Electronics should have the highest revenue → should be first row when ordered DESC164    first_cat = str(rows[0].get(cat_col, "")).lower() if rows else ""165    correct_order = first_cat == "electronics"166 167    all_three = cat_coverage >= 0.99168 169    if all_three and correct_order:170        return 1.0, "Excellent! All 3 categories with correct revenue, ordered highest first."171    elif all_three:172        return 0.75, (173            "All 3 categories present but not sorted highest-to-lowest. "174            "Add ORDER BY total_revenue DESC."175        )176    elif cat_coverage >= 0.66:177        return 0.5, (178            f"Found {int(round(cat_coverage * 3))}/3 categories. "179            "Check your JOIN condition and GROUP BY clause."180        )181    else:182        return 0.2, (183            "Missing categories. Join order_items with products on product_id, "184            "then GROUP BY products.category."185        )186 187 188# ---------------------------------------------------------------------------189# Grader: window_ranking190# ---------------------------------------------------------------------------191 192def grade_window_ranking(193    rows: List[dict],194    columns: List[str],195    error: Optional[str],196) -> Tuple[float, str]:197    """198    Expected: all customers who have placed orders, with:199      - customer name200      - most recent order date201      - rank by total spending (rank 1 = highest)202    Results ordered by rank ascending.203    """204    if error:205        return 0.0, f"Query error: {error}"206    if not rows:207        return 0.0, "No rows returned. Try aggregating orders per customer."208 209    name_col = next((c for c in columns if "name" in c.lower()), None)210    rank_col = next((c for c in columns if "rank" in c.lower()), None)211    date_col = next(212        (c for c in columns if "date" in c.lower() or "recent" in c.lower()),213        None,214    )215 216    score = 0.0217    feedback_parts: List[str] = []218 219    if name_col:220        score += 0.15221        feedback_parts.append("customer names present")222 223    if rank_col:224        ranks = []225        for r in rows:226            v = r.get(rank_col)227            if v is not None:228                try:229                    ranks.append(int(v))230                except (ValueError, TypeError):231                    pass232        if 1 in ranks:233            score += 0.35234            feedback_parts.append("rank column present with rank 1")235        else:236            score += 0.15237            feedback_parts.append("rank column present but rank 1 missing")238 239    if date_col:240        score += 0.2241        feedback_parts.append("order date column present")242 243    # Rows should be ordered by rank ascending244    if rank_col and len(rows) > 1:245        rank_vals = []246        for r in rows:247            try:248                rank_vals.append(int(r.get(rank_col, 999)))249            except (ValueError, TypeError):250                pass251        if rank_vals and rank_vals == sorted(rank_vals):252            score += 0.3253            feedback_parts.append("results ordered by rank correctly")254 255    score = min(1.0, score)256 257    if score >= 0.9:258        return 1.0, "Excellent! Correct ranking with proper ordering by rank."259    elif score >= 0.5:260        return round(score, 2), "Good progress: " + ", ".join(feedback_parts) + "."261    elif score > 0:262        return round(score, 2), (263            "Partial credit: " + ", ".join(feedback_parts) + ". Keep refining."264        )265    else:266        return 0.1, (267            "Incorrect. Use RANK() OVER (ORDER BY total_spent DESC) or a correlated "268            "subquery to compute ranks. Include name, most recent date, and rank."269        )270