Hariprita/nl2sql-openenv
0
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 