Tsah00/sql-env
0
1"""2server/tasks.py - Task definitions and automated graders for the SQL Environment.3 4Defines 3 difficulty tiers (easy, medium, hard) each with 3 tasks.5Each task includes:6 - A natural language description7 - A reference solution SQL query8 - An automated grader function9 - A reward calculator10 11Database schema (e-commerce):12 customers (id, name, email, city, country, created_at)13 products (id, name, category, price, stock)14 orders (id, customer_id, order_date, total_amount, status)15 order_items(id, order_id, product_id, quantity, unit_price)16"""17 18from __future__ import annotations19 20import sqlite321import re22from typing import Any, Dict, List, Optional, Tuple23 24 25# ---------------------------------------------------------------------------26# Schema description returned to the agent27# ---------------------------------------------------------------------------28 29SCHEMA_INFO = """30DATABASE SCHEMA (SQLite):31 32TABLE customers:33 id INTEGER PRIMARY KEY34 name TEXT NOT NULL35 email TEXT UNIQUE NOT NULL36 city TEXT37 country TEXT38 created_at TEXT -- ISO date e.g. '2023-01-15'39 40TABLE products:41 id INTEGER PRIMARY KEY42 name TEXT NOT NULL43 category TEXT NOT NULL44 price REAL NOT NULL45 stock INTEGER DEFAULT 046 47TABLE orders:48 id INTEGER PRIMARY KEY49 customer_id INTEGER REFERENCES customers(id)50 order_date TEXT -- ISO date e.g. '2023-06-20'51 total_amount REAL NOT NULL52 status TEXT -- 'completed', 'pending', 'cancelled'53 54TABLE order_items:55 id INTEGER PRIMARY KEY56 order_id INTEGER REFERENCES orders(id)57 product_id INTEGER REFERENCES products(id)58 quantity INTEGER NOT NULL59 unit_price REAL NOT NULL60 61SAMPLE DATA HINTS:62 - 20 customers across 9 countries63 - 20 products across 5 categories (Electronics, Clothing, Sports, Home, Accessories, Stationery)64 - 30 orders spanning 2023-2024, statuses: completed / pending / cancelled65 - 50 order_items linking orders to products66 67ROLE CONTEXT:68 You are a data analyst at an e-commerce company. Business stakeholders69 (marketing, finance, CRM, merchandising) submit ad-hoc data requests that70 you fulfil by writing SQL queries against this database.71""".strip()72 73 74# ---------------------------------------------------------------------------75# Database seeding - deterministic test data76# ---------------------------------------------------------------------------77 78def seed_database(conn: sqlite3.Connection) -> None:79 """Populate the database with deterministic e-commerce test data."""80 conn.executescript("""81 CREATE TABLE IF NOT EXISTS customers (82 id INTEGER PRIMARY KEY,83 name TEXT NOT NULL,84 email TEXT UNIQUE NOT NULL,85 city TEXT,86 country TEXT,87 created_at TEXT88 );89 90 CREATE TABLE IF NOT EXISTS products (91 id INTEGER PRIMARY KEY,92 name TEXT NOT NULL,93 category TEXT NOT NULL,94 price REAL NOT NULL,95 stock INTEGER DEFAULT 096 );97 98 CREATE TABLE IF NOT EXISTS orders (99 id INTEGER PRIMARY KEY,100 customer_id INTEGER REFERENCES customers(id),101 order_date TEXT,102 total_amount REAL NOT NULL,103 status TEXT104 );105 106 CREATE TABLE IF NOT EXISTS order_items (107 id INTEGER PRIMARY KEY,108 order_id INTEGER REFERENCES orders(id),109 product_id INTEGER REFERENCES products(id),110 quantity INTEGER NOT NULL,111 unit_price REAL NOT NULL112 );113 """)114 115 # Customers116 customers = [117 (1, "Alice Johnson", "alice@example.com", "New York", "USA", "2022-01-10"),118 (2, "Bob Smith", "bob@example.com", "London", "UK", "2022-02-15"),119 (3, "Carol White", "carol@example.com", "Toronto", "Canada", "2022-03-20"),120 (4, "David Brown", "david@example.com", "Sydney", "Australia","2022-04-05"),121 (5, "Eva Martinez", "eva@example.com", "Madrid", "Spain", "2022-05-12"),122 (6, "Frank Lee", "frank@example.com", "Tokyo", "Japan", "2022-06-18"),123 (7, "Grace Kim", "grace@example.com", "Seoul", "Korea", "2022-07-22"),124 (8, "Henry Wang", "henry@example.com", "Beijing", "China", "2022-08-30"),125 (9, "Isabel Garcia", "isabel@example.com", "Mexico City", "Mexico", "2022-09-14"),126 (10, "Jack Taylor", "jack@example.com", "Chicago", "USA", "2022-10-01"),127 (11, "Kate Davis", "kate@example.com", "Los Angeles", "USA", "2022-11-11"),128 (12, "Liam Wilson", "liam@example.com", "Manchester", "UK", "2022-12-05"),129 (13, "Mia Anderson", "mia@example.com", "Vancouver", "Canada", "2023-01-08"),130 (14, "Noah Thomas", "noah@example.com", "Melbourne", "Australia","2023-02-14"),131 (15, "Olivia Harris", "olivia@example.com", "Barcelona", "Spain", "2023-03-21"),132 (16, "Paul Martin", "paul@example.com", "Osaka", "Japan", "2023-04-17"),133 (17, "Quinn Lewis", "quinn@example.com", "Busan", "Korea", "2023-05-09"),134 (18, "Rachel Walker", "rachel@example.com", "Shanghai", "China", "2023-06-25"),135 (19, "Sam Hall", "sam@example.com", "Guadalajara", "Mexico", "2023-07-30"),136 (20, "Tina Young", "tina@example.com", "Houston", "USA", "2023-08-16"),137 ]138 conn.executemany(139 "INSERT OR IGNORE INTO customers VALUES (?,?,?,?,?,?)", customers140 )141 142 # Products143 products = [144 (1, "Laptop Pro", "Electronics", 1299.99, 50),145 (2, "Wireless Mouse", "Electronics", 29.99, 200),146 (3, "USB-C Hub", "Electronics", 49.99, 150),147 (4, "Mechanical Keyboard","Electronics", 89.99, 100),148 (5, "4K Monitor", "Electronics", 399.99, 40),149 (6, "Running Shoes", "Clothing", 119.99, 80),150 (7, "Yoga Mat", "Sports", 34.99, 120),151 (8, "Dumbbell Set", "Sports", 79.99, 60),152 (9, "Water Bottle", "Sports", 19.99, 300),153 (10, "Backpack", "Accessories", 59.99, 90),154 (11, "Sunglasses", "Accessories", 79.99, 70),155 (12, "Coffee Maker", "Home", 149.99, 45),156 (13, "Air Purifier", "Home", 199.99, 30),157 (14, "Desk Lamp", "Home", 39.99, 110),158 (15, "Notebook (set of 3)","Stationery", 14.99, 250),159 (16, "Ballpoint Pens", "Stationery", 9.99, 400),160 (17, "Webcam HD", "Electronics", 69.99, 85),161 (18, "Headphones BT", "Electronics", 129.99, 65),162 (19, "Resistance Bands", "Sports", 24.99, 180),163 (20, "Throw Pillow", "Home", 29.99, 95),164 ]165 conn.executemany(166 "INSERT OR IGNORE INTO products VALUES (?,?,?,?,?)", products167 )168 169 # Orders170 orders = [171 (1, 1, "2023-01-15", 1329.98, "completed"),172 (2, 2, "2023-01-20", 59.98, "completed"),173 (3, 3, "2023-02-05", 399.99, "completed"),174 (4, 4, "2023-02-14", 89.99, "pending"),175 (5, 5, "2023-03-10", 259.98, "completed"),176 (6, 6, "2023-03-22", 149.99, "completed"),177 (7, 7, "2023-04-08", 54.98, "cancelled"),178 (8, 8, "2023-04-19", 479.97, "completed"),179 (9, 9, "2023-05-02", 79.99, "completed"),180 (10, 10, "2023-05-17", 209.97, "completed"),181 (11, 11, "2023-06-03", 129.99, "completed"),182 (12, 12, "2023-06-25", 39.99, "pending"),183 (13, 13, "2023-07-11", 229.97, "completed"),184 (14, 14, "2023-07-28", 599.98, "completed"),185 (15, 15, "2023-08-15", 44.98, "completed"),186 (16, 16, "2023-09-01", 349.97, "completed"),187 (17, 17, "2023-09-18", 24.99, "cancelled"),188 (18, 18, "2023-10-05", 279.98, "completed"),189 (19, 19, "2023-10-22", 89.97, "completed"),190 (20, 20, "2023-11-08", 159.99, "completed"),191 (21, 1, "2023-11-25", 1399.98, "completed"),192 (22, 2, "2023-12-10", 199.99, "completed"),193 (23, 3, "2024-01-05", 99.98, "completed"),194 (24, 4, "2024-01-20", 449.97, "pending"),195 (25, 5, "2024-02-14", 259.98, "completed"),196 (26, 10, "2024-02-28", 129.99, "completed"),197 (27, 11, "2024-03-15", 349.98, "completed"),198 (28, 1, "2024-03-28", 79.99, "completed"),199 (29, 2, "2024-04-10", 59.99, "completed"),200 (30, 3, "2024-04-22", 199.98, "completed"),201 ]202 conn.executemany(203 "INSERT OR IGNORE INTO orders VALUES (?,?,?,?,?)", orders204 )205 206 # Order items207 order_items = [208 (1, 1, 1, 1, 1299.99),209 (2, 1, 2, 1, 29.99),210 (3, 2, 2, 1, 29.99),211 (4, 2, 3, 1, 49.99),212 (5, 3, 5, 1, 399.99),213 (6, 4, 4, 1, 89.99),214 (7, 5, 6, 1, 119.99),215 (8, 5, 7, 2, 34.99),216 (9, 6, 12, 1, 149.99),217 (10, 7, 7, 1, 34.99),218 (11, 7, 9, 1, 19.99),219 (12, 8, 5, 1, 399.99),220 (13, 8, 2, 2, 29.99),221 (14, 8, 3, 1, 49.99),222 (15, 9, 8, 1, 79.99),223 (16, 10, 18, 1, 129.99),224 (17, 10, 2, 2, 29.99),225 (18, 10, 14, 1, 39.99),226 (19, 11, 18, 1, 129.99),227 (20, 12, 14, 1, 39.99),228 (21, 13, 6, 1, 119.99),229 (22, 13, 9, 2, 19.99),230 (23, 13, 15, 6, 14.99),231 (24, 14, 1, 1, 1299.99),232 (25, 14, 16, 4, 9.99), # 4*9.99=39.99 but order total 599.98 -- we simplify233 (26, 15, 7, 1, 34.99),234 (27, 15, 9, 1, 19.99),235 (28, 16, 5, 1, 399.99), # corrected to match total approx236 (29, 16, 2, 2, 29.99),237 (30, 17, 19, 1, 24.99),238 (31, 18, 1, 1, 279.98), # simplified239 (32, 19, 9, 3, 19.99),240 (33, 19, 16, 4, 9.99),241 (34, 20, 13, 1, 159.99), # simplified242 (35, 21, 1, 1, 1299.99),243 (36, 21, 4, 1, 89.99),244 (37, 22, 13, 1, 199.99),245 (38, 23, 4, 1, 89.99),246 (39, 23, 9, 1, 19.99),247 (40, 24, 5, 1, 399.99),248 (41, 24, 18, 1, 129.99),249 (42, 25, 6, 1, 119.99),250 (43, 25, 7, 2, 34.99),251 (44, 26, 18, 1, 129.99),252 (45, 27, 5, 1, 399.99), # simplified253 (46, 27, 2, 2, 29.99), # simplified254 (47, 28, 8, 1, 79.99),255 (48, 29, 16, 6, 9.99),256 (49, 30, 12, 1, 149.99),257 (50, 30, 7, 1, 34.99), # simplified258 ]259 conn.executemany(260 "INSERT OR IGNORE INTO order_items VALUES (?,?,?,?,?)", order_items261 )262 conn.commit()263 264 265# ---------------------------------------------------------------------------266# Grader helpers267# ---------------------------------------------------------------------------268 269def _normalize_rows(rows: List[Dict]) -> List[Dict]:270 """Round floats to 2 decimal places for comparison."""271 normalized = []272 for row in rows:273 norm = {}274 for k, v in row.items():275 if isinstance(v, float):276 norm[k] = round(v, 2)277 else:278 norm[k] = v279 normalized.append(norm)280 return normalized281 282 283def _row_to_values(row: Dict) -> tuple:284 """Extract values in column-name order (for alias-agnostic comparison)."""285 return tuple(row[k] for k in sorted(row.keys()))286 287 288def _rows_match(agent_rows: List[Dict], expected_rows: List[Dict],289 ordered: bool = False) -> float:290 """291 Return correctness score 0.0-1.0 comparing agent vs expected results.292 Uses set comparison for unordered, sequence comparison for ordered.293 Partial credit: ratio of matched rows.294 295 Column alias normalization: if the agent uses different column names but296 returns the same values (e.g. 'total' instead of 'total_spent'), the score297 is taken as the max of key-match and value-only-match so cosmetic aliases298 are not penalised.299 """300 if not expected_rows:301 return 1.0 if not agent_rows else 0.0302 303 agent_norm = _normalize_rows(agent_rows)304 expected_norm = _normalize_rows(expected_rows)305 306 def _multiset_jaccard(a_tuples: list, e_tuples: list) -> float:307 """308 Jaccard-style multiset score: matched / (|A| + |E| - matched).309 Penalises both missing rows (low recall) and extra rows (low precision).310 Prevents the exploit of returning the entire table to score 1.0.311 """312 a_counter: Dict = {}313 for r in a_tuples:314 a_counter[r] = a_counter.get(r, 0) + 1315 e_counter: Dict = {}316 for r in e_tuples:317 e_counter[r] = e_counter.get(r, 0) + 1318 matched = sum(min(cnt, a_counter.get(r, 0)) for r, cnt in e_counter.items())319 union = len(a_tuples) + len(e_tuples) - matched320 return matched / union if union else 1.0321 322 def _score_with_keys(a_rows, e_rows):323 if ordered:324 if not a_rows:325 return 0.0326 matches = sum(1 for a, e in zip(a_rows, e_rows) if a == e)327 # Penalise wrong length: score against the longer of the two328 denom = max(len(a_rows), len(e_rows))329 return matches / denom330 return _multiset_jaccard(331 [tuple(sorted(r.items())) for r in a_rows],332 [tuple(sorted(r.items())) for r in e_rows],333 )334 335 def _score_values_only(a_rows, e_rows):336 """Compare only values (sorted by col name), ignoring column aliases."""337 if not a_rows or len(a_rows[0]) != len(e_rows[0]):338 return 0.0339 a_vals = [_row_to_values(r) for r in a_rows]340 e_vals = [_row_to_values(r) for r in e_rows]341 if ordered:342 matches = sum(1 for a, e in zip(a_vals, e_vals) if a == e)343 denom = max(len(a_vals), len(e_vals))344 return matches / denom345 return _multiset_jaccard(a_vals, e_vals)346 347 key_score = _score_with_keys(agent_norm, expected_norm)348 val_score = _score_values_only(agent_norm, expected_norm)349 return max(key_score, val_score)350 351 352def _query_complexity_penalty(query: str) -> float:353 """354 Return efficiency bonus (0.0-0.2) based on query simplicity.355 Penalises unnecessary SELECT *, excessive subqueries, CROSS JOINs.356 """357 q = query.upper()358 penalty = 0.0359 if "SELECT *" in q:360 penalty += 0.05361 if q.count("SELECT") > 3:362 penalty += 0.05363 if "CROSS JOIN" in q:364 penalty += 0.1365 return max(0.0, 0.2 - penalty)366 367 368def _has_required_keywords(query: str, keywords: List[str]) -> bool:369 q = query.upper()370 return all(kw.upper() in q for kw in keywords)371 372 373# ---------------------------------------------------------------------------374# Task registry375# ---------------------------------------------------------------------------376 377TASKS: Dict[str, Dict] = {}378 379 380def _register(task_id: str, difficulty: str, description: str,381 reference_sql: str, required_keywords: List[str],382 expected_columns: List[str], ordered: bool = False):383 TASKS[task_id] = {384 "id": task_id,385 "difficulty": difficulty,386 "description": description,387 "reference_sql": reference_sql,388 "required_keywords": required_keywords,389 "expected_columns": expected_columns,390 "ordered": ordered,391 }392 393 394# ── EASY ────────────────────────────────────────────────────────────────────395 396_register(397 task_id="easy_1",398 difficulty="easy",399 description=(400 "The marketing team is launching a US-only promotional email campaign. "401 "Retrieve the names and emails of all customers from the USA. "402 "Return columns: name, email."403 ),404 reference_sql="SELECT name, email FROM customers WHERE country = 'USA'",405 required_keywords=["SELECT", "FROM", "WHERE"],406 expected_columns=["name", "email"],407)408 409_register(410 task_id="easy_2",411 difficulty="easy",412 description=(413 "The finance team needs a daily fulfilment report. "414 "Count the total number of orders with status 'completed'. "415 "Return a single column named: total_completed."416 ),417 reference_sql=(418 "SELECT COUNT(*) AS total_completed FROM orders WHERE status = 'completed'"419 ),420 required_keywords=["SELECT", "COUNT", "WHERE"],421 expected_columns=["total_completed"],422)423 424_register(425 task_id="easy_3",426 difficulty="easy",427 description=(428 "The merchandising team is building a premium product catalogue. "429 "List the top 5 most expensive products by price, highest first. "430 "Return columns: name, category, price."431 ),432 reference_sql=(433 "SELECT name, category, price FROM products "434 "ORDER BY price DESC LIMIT 5"435 ),436 required_keywords=["SELECT", "ORDER BY", "LIMIT"],437 expected_columns=["name", "category", "price"],438 ordered=True,439)440 441# ── MEDIUM ───────────────────────────────────────────────────────────────────442 443_register(444 task_id="medium_1",445 difficulty="medium",446 description=(447 "The loyalty team wants to identify top spenders for a VIP rewards programme. "448 "For each customer, calculate their total spending across all completed orders. "449 "Include only customers who have placed at least one completed order. "450 "Return columns: name, total_spent. Order by total_spent descending."451 ),452 reference_sql="""453 SELECT c.name, SUM(o.total_amount) AS total_spent454 FROM customers c455 JOIN orders o ON c.id = o.customer_id456 WHERE o.status = 'completed'457 GROUP BY c.id, c.name458 ORDER BY total_spent DESC459 """.strip(),460 required_keywords=["SELECT", "JOIN", "GROUP BY", "SUM"],461 expected_columns=["name", "total_spent"],462 ordered=True,463)464 465_register(466 task_id="medium_2",467 difficulty="medium",468 description=(469 "The inventory team needs to identify dead-stock items for a clearance sale. "470 "Find all products that have NEVER appeared in any order. "471 "Return columns: name, category, price."472 ),473 reference_sql="""474 SELECT p.name, p.category, p.price475 FROM products p476 LEFT JOIN order_items oi ON p.id = oi.product_id477 WHERE oi.id IS NULL478 """.strip(),479 required_keywords=["SELECT", "LEFT JOIN", "WHERE"],480 expected_columns=["name", "category", "price"],481)482 483_register(484 task_id="medium_3",485 difficulty="medium",486 description=(487 "The finance team is building a monthly revenue trend report for 2023. "488 "Calculate the average order value for each month in 2023. "489 "Format the month as 'YYYY-MM'. "490 "Return columns: month, avg_order_value. Order by month ascending."491 ),492 reference_sql="""493 SELECT STRFTIME('%Y-%m', order_date) AS month,494 ROUND(AVG(total_amount), 2) AS avg_order_value495 FROM orders496 WHERE order_date LIKE '2023%'497 GROUP BY month498 ORDER BY month ASC499 """.strip(),500 required_keywords=["SELECT", "AVG", "GROUP BY"],501 expected_columns=["month", "avg_order_value"],502 ordered=True,503)504 505# ── HARD ──────────────────────────────────────────────────────────────────────506 507_register(508 task_id="hard_1",509 difficulty="hard",510 description=(511 "The CRM team wants to segment high-value customers for a premium tier. "512 "Find all customers whose total completed spending is above the average "513 "total spending across all customers who have made at least one completed order. "514 "Use a CTE for clarity. "515 "Return columns: name, total_spent. Order by total_spent descending."516 ),517 reference_sql="""518 WITH customer_totals AS (519 SELECT c.name, SUM(o.total_amount) AS total_spent520 FROM customers c521 JOIN orders o ON c.id = o.customer_id522 WHERE o.status = 'completed'523 GROUP BY c.id, c.name524 )525 SELECT name, total_spent526 FROM customer_totals527 WHERE total_spent > (SELECT AVG(total_spent) FROM customer_totals)528 ORDER BY total_spent DESC529 """.strip(),530 required_keywords=["SELECT", "GROUP BY", "AVG"],531 expected_columns=["name", "total_spent"],532 ordered=True,533)534 535_register(536 task_id="hard_2",537 difficulty="hard",538 description=(539 "The category management team needs a 'category hero' SKU report for the homepage. "540 "For each product category, find the best-selling product by total quantity sold. "541 "If two products tie, return both (use RANK, not ROW_NUMBER). "542 "Return columns: category, product_name, total_quantity."543 ),544 reference_sql="""545 WITH product_sales AS (546 SELECT p.category,547 p.name AS product_name,548 SUM(oi.quantity) AS total_quantity549 FROM products p550 JOIN order_items oi ON p.id = oi.product_id551 GROUP BY p.id, p.category, p.name552 ),553 ranked AS (554 SELECT category, product_name, total_quantity,555 RANK() OVER (PARTITION BY category ORDER BY total_quantity DESC) AS rnk556 FROM product_sales557 )558 SELECT category, product_name, total_quantity559 FROM ranked560 WHERE rnk = 1561 ORDER BY category, product_name562 """.strip(),563 required_keywords=["SELECT", "GROUP BY", "JOIN"],564 expected_columns=["category", "product_name", "total_quantity"],565)566 567_register(568 task_id="hard_3",569 difficulty="hard",570 description=(571 "The retention team wants to reward customers who have been active every "572 "year since the platform launched. Find customers who placed at least one "573 "order in each of the years 2022, 2023, AND 2024 (all three). "574 "Return columns: name, email. Order by name ascending."575 ),576 reference_sql="""577 SELECT c.name, c.email578 FROM customers c579 WHERE (580 SELECT COUNT(DISTINCT STRFTIME('%Y', o.order_date))581 FROM orders o582 WHERE o.customer_id = c.id583 AND STRFTIME('%Y', o.order_date) IN ('2022', '2023', '2024')584 ) = 3585 ORDER BY c.name ASC586 """.strip(),587 required_keywords=["SELECT", "WHERE"],588 expected_columns=["name", "email"],589 ordered=True,590)591 592 593# ---------------------------------------------------------------------------594# Grader595# ---------------------------------------------------------------------------596 597def grade(598 task_id: str,599 agent_query: str,600 conn: sqlite3.Connection,601) -> Tuple[float, str, List[Dict], List[Dict]]:602 """603 Grade an agent's SQL query against the reference solution.604 605 Returns:606 (reward, message, agent_rows, expected_rows)607 reward: float 0.0-1.0608 message: human-readable feedback609 agent_rows: what the agent's query returned610 expected_rows: what the reference returned611 """612 # Reward is always in the open interval (0.0, 1.0) — never exactly 0 or 1.613 # This satisfies Phase 2 validation which requires scores strictly within (0, 1).614 _REWARD_MIN = 0.01615 _REWARD_MAX = 0.99616 617 task = TASKS.get(task_id)618 if task is None:619 return _REWARD_MIN, f"Unknown task_id: {task_id}", [], []620 621 # Run reference solution622 try:623 cur = conn.execute(task["reference_sql"])624 cols = [d[0] for d in cur.description]625 expected_rows = [dict(zip(cols, row)) for row in cur.fetchall()]626 except Exception as e:627 return _REWARD_MIN, f"Reference SQL failed (bug): {e}", [], []628 629 # Run agent query630 try:631 cur = conn.execute(agent_query)632 cols = [d[0] for d in cur.description]633 agent_rows = [dict(zip(cols, row)) for row in cur.fetchall()]634 except Exception as e:635 return _REWARD_MIN, f"Query error: {e}", [], expected_rows636 637 # Score correctness638 correctness = _rows_match(agent_rows, expected_rows, task["ordered"])639 640 # Keyword bonus (checks structural correctness)641 kw_bonus = 0.1 if _has_required_keywords(agent_query, task["required_keywords"]) else 0.0642 643 # Efficiency bonus644 efficiency = _query_complexity_penalty(agent_query)645 646 # Total reward — clamped to open interval (0.01, 0.99)647 raw = correctness * 0.7 + kw_bonus + efficiency648 reward = round(min(_REWARD_MAX, max(_REWARD_MIN, raw)), 4)649 650 # Build message651 if correctness >= 0.99:652 msg = "Correct! Full marks for result accuracy."653 elif correctness >= 0.5:654 msg = f"Partially correct ({correctness:.0%} rows matched)."655 else:656 msg = f"Incorrect results ({correctness:.0%} rows matched)."657 658 if reward >= _REWARD_MAX:659 msg += " Near-perfect score!"660 661 return reward, msg, agent_rows, expected_rows662 663 664def get_task_by_difficulty(difficulty: str) -> Optional[Dict]:665 """Return the first task matching the given difficulty."""666 for task in TASKS.values():667 if task["difficulty"] == difficulty:668 return task669 return None670 671 672def get_all_tasks_by_difficulty(difficulty: str) -> List[Dict]:673 return [t for t in TASKS.values() if t["difficulty"] == difficulty]674 