{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {
    "id": "f1pjVyZ02Gs6"
   },
   "source": [
    "# 11-766 HW3 Question 2: Making MCP Tool Calls with a Language Model\n",
    "\n",
    "In this assignment, you will:\n",
    "1. Measure how often a small model fails to produce reliable structured outputs.\n",
    "2. Compare prompt-only extraction against schema-constrained generation.\n",
    "3. Stress-test the system with edge cases that try to break JSON or schema validity.\n",
    "\n",
    "**Runtime**: This notebook requires a GPU. Use Colab with at least a T4 or L4 GPU.\n",
    "\n",
    "**Important note**: exact outputs may vary slightly across machines and library versions. You are not graded on reproducing one exact table of numbers or one particular unicorn; you are graded on whether your analysis correctly explains the structural and semantic failure modes you observe."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "id": "K9sm2DknBD7I"
   },
   "source": [
    "# 0. Setup"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "colab": {
     "base_uri": "https://localhost:8080/"
    },
    "id": "RY0lmNKZ-_1-",
    "outputId": "c322204c-f7a6-41f6-dd33-2e6897aee6e8"
   },
   "outputs": [],
   "source": [
    "!pip install fastmcp\n",
    "!pip install cairosvg"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "id": "_WYgad2_KoxW"
   },
   "outputs": [],
   "source": [
    "import math\n",
    "import random\n",
    "import json\n",
    "import ast\n",
    "import xml.etree.ElementTree as ET\n",
    "from dataclasses import dataclass, field\n",
    "from typing import Dict, List, Any, Tuple\n",
    "from IPython.display import SVG, display\n",
    "from contextlib import asynccontextmanager\n",
    "from fastmcp import FastMCP, Client\n",
    "from fastmcp.server.dependencies import get_context\n",
    "from openai import OpenAI\n",
    "from google.colab import userdata\n",
    "import base64\n",
    "import cairosvg\n",
    "\n",
    "import warnings\n",
    "warnings.filterwarnings(\"ignore\", category=DeprecationWarning)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "id": "n9mNZIJAh3kv"
   },
   "source": [
    "SVG creation utilities\n",
    "\n",
    "*You should not need to modify this codeblock.*"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "id": "7C6BZGPDLPhZ"
   },
   "outputs": [],
   "source": [
    "@dataclass\n",
    "class SVGElement:\n",
    "    tag: str\n",
    "    attrs: Dict[str, Any]\n",
    "\n",
    "@dataclass\n",
    "class SVGDoc:\n",
    "    \"\"\"A class to represent and build an SVG document. This is used by the server.\"\"\"\n",
    "    width: int = None\n",
    "    height: int = None\n",
    "    elements: Dict[str, SVGElement] = field(default_factory=dict)\n",
    "\n",
    "    def create_svg(self, width: int, height: int):\n",
    "        \"\"\"Initializes the SVG canvas with the given width and height.\"\"\"\n",
    "        self.width = width\n",
    "        self.height = height\n",
    "\n",
    "    def add_element(self, tag: str, attrs: Dict[str, Any], elem_id: str):\n",
    "        \"\"\"Adds a new SVG element to the document.\"\"\"\n",
    "        if elem_id in self.elements:\n",
    "            raise ValueError(f\"Duplicate id: {elem_id}\")\n",
    "        attrs = dict(attrs)\n",
    "        attrs[\"id\"] = elem_id\n",
    "        self.elements[elem_id] = SVGElement(tag=tag, attrs=attrs)\n",
    "\n",
    "    def set_attr(self, elem_id: str, key: str, value: Any):\n",
    "        \"\"\"Sets an attribute for an existing SVG element.\"\"\"\n",
    "        if elem_id not in self.elements:\n",
    "            raise ValueError(f\"Unknown id: {elem_id}\")\n",
    "        self.elements[elem_id].attrs[key] = value\n",
    "\n",
    "    def list_ids(self):\n",
    "        \"\"\"Lists all element IDs currently in the document.\"\"\"\n",
    "        return list(self.elements.keys())\n",
    "\n",
    "    def export_svg(self) -> str:\n",
    "        \"\"\"Exports the current SVG document as an XML string.\"\"\"\n",
    "        if self.width is None or self.height is None:\n",
    "            raise ValueError(\"SVG canvas not created\")\n",
    "\n",
    "        root = ET.Element(\n",
    "            \"svg\",\n",
    "            {\n",
    "                \"xmlns\": \"http://www.w3.org/2000/svg\",\n",
    "                \"width\": str(self.width),\n",
    "                \"height\": str(self.height),\n",
    "                \"viewBox\": f\"0 0 {self.width} {self.height}\"\n",
    "            }\n",
    "        )\n",
    "\n",
    "        for elem in self.elements.values():\n",
    "            attrib = {k: str(v) for k, v in elem.attrs.items()}\n",
    "            ET.SubElement(root, elem.tag, attrib)\n",
    "\n",
    "        return ET.tostring(root, encoding=\"unicode\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "id": "fc3xB0R-iPBA"
   },
   "source": [
    "MCP functions and helpers\n",
    "\n",
    "*You should not need to modify this codeblock.*"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "id": "vV7maNK0Lzf2"
   },
   "outputs": [],
   "source": [
    "# --- Shared state ---\n",
    "@dataclass\n",
    "class SVGState:\n",
    "    canvas: Dict[str, int] = field(default_factory=dict)\n",
    "    elements: Dict[str, Any] = field(default_factory=dict)\n",
    "\n",
    "@asynccontextmanager\n",
    "async def lifespan(server: FastMCP):\n",
    "    yield {\"state\": SVGState()}\n",
    "\n",
    "_default_state = SVGState()\n",
    "\n",
    "# --- Create MCP server ---\n",
    "mcp = FastMCP(\"svg-creator\", lifespan=lifespan)\n",
    "\n",
    "# --- Helper functions ---\n",
    "def get_canvas_and_elements() -> SVGState:\n",
    "    \"\"\"Helper function that gets the canvas state..\"\"\"\n",
    "    try:\n",
    "        ctx = get_context()\n",
    "        state = ctx.request_context.lifespan_context[\"state\"]\n",
    "    except RuntimeError:\n",
    "        # No active MCP context. This fallback will be used when you are\n",
    "        # calling these functions directly (Question 2.1), rather than having\n",
    "        # the LLM agent call them.\n",
    "        state = _default_state\n",
    "    return state.canvas, state.elements\n",
    "\n",
    "def get_elements_dict() -> SVGState:\n",
    "    \"\"\"Helper function that gets the elements_dict state..\"\"\"\n",
    "    ctx = get_context()\n",
    "    return ctx.request_context.lifespan_context[\"state\"].elements_dict\n",
    "\n",
    "@mcp.tool()\n",
    "def create_svg(width: int, height: int) -> str:\n",
    "    \"\"\"Creates a new SVG canvas with the specified width and height.\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "\n",
    "    canvas[\"width\"] = width\n",
    "    canvas[\"height\"] = height\n",
    "    elements.clear()\n",
    "    return f\"Canvas created: {width}x{height}\"\n",
    "\n",
    "def render_in_colab(svg_string):\n",
    "    \"\"\"Renders the current SVG document directly in a Colab output cell.\"\"\"\n",
    "    display(SVG(svg_string))\n",
    "\n",
    "# --- MCP functions ---\n",
    "@mcp.tool()\n",
    "def add_circle(cx: int, cy: int, r: int, id: str) -> str:\n",
    "    \"\"\"Adds a circle. cx/cy = center, r = radius, id = unique element ID.\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "    suffix = \"replaced previous version\" if id in elements else \"added\"\n",
    "    elements[id] = SVGElement(tag=\"circle\", attrs={\"cx\": cx, \"cy\": cy, \"r\": r, \"id\": id})\n",
    "    return f\"Circle '{id}' {suffix}\"\n",
    "\n",
    "@mcp.tool()\n",
    "def add_ellipse(cx: int, cy: int, rx: int, ry: int, id: str) -> str:\n",
    "    \"\"\"Adds an ellipse. cx/cy = center, rx/ry = radii, id = unique element ID.\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "    suffix = \"replaced previous version\" if id in elements else \"added\"\n",
    "    elements[id] = SVGElement(tag=\"ellipse\", attrs={\"cx\": cx, \"cy\": cy, \"rx\": rx, \"ry\": ry, \"id\": id})\n",
    "    return f\"Ellipse '{id}' {suffix}\"\n",
    "\n",
    "@mcp.tool()\n",
    "def add_rect(x: int, y: int, width: int, height: int, id: str) -> str:\n",
    "    \"\"\"Adds a rectangle. x/y = top-left corner, id = unique element ID.\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "    suffix = \"replaced previous version\" if id in elements else \"added\"\n",
    "    elements[id] = SVGElement(tag=\"rect\", attrs={\"x\": x, \"y\": y, \"width\": width, \"height\": height, \"id\": id})\n",
    "    return f\"Rectangle '{id}' {suffix}\"\n",
    "\n",
    "@mcp.tool()\n",
    "def add_triangle(points: List[List[int]], id: str) -> str:\n",
    "    \"\"\"Adds a triangle polygon. points = [[x1,y1],[x2,y2],[x3,y3]], id = unique element ID.\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "    suffix = \"replaced previous version\" if id in elements else \"added\"\n",
    "    points_str = \" \".join(f\"{x},{y}\" for x, y in points)\n",
    "    elements[id] = SVGElement(tag=\"polygon\", attrs={\"points\": points_str, \"id\": id})\n",
    "    return f\"Triangle '{id}' {suffix}\"\n",
    "\n",
    "@mcp.tool()\n",
    "def add_polyline(points: List[List[int]], id: str) -> str:\n",
    "    \"\"\"Adds a polyline (open path). points = [[x1,y1],...], id = unique element ID.\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "    suffix = \"replaced previous version\" if id in elements else \"added\"\n",
    "    points_str = \" \".join(f\"{x},{y}\" for x, y in points)\n",
    "    elements[id] = SVGElement(tag=\"polyline\", attrs={\"points\": points_str, \"fill\": \"none\", \"id\": id})\n",
    "    return f\"Polyline '{id}' {suffix}\"\n",
    "\n",
    "@mcp.tool()\n",
    "def delete_shape(id:str) -> str:\n",
    "    \"\"\"Removes the shape with the specified ID from the canvas. Use this to try again.\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "    del elements[id]\n",
    "    return f\"Shape '{id}' deleted from the canvas.\"\n",
    "\n",
    "@mcp.tool()\n",
    "def set_fill(id: str, color: str) -> str:\n",
    "    \"\"\"Sets the fill color of an element. color can be a name ('red') or hex ('#ff0000').\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "    if id not in elements:\n",
    "        raise ValueError(f\"Unknown id: {id}\")\n",
    "    elements[id].attrs[\"fill\"] = color\n",
    "    return f\"Fill of '{id}' set to {color}\"\n",
    "\n",
    "@mcp.tool()\n",
    "def set_stroke(id: str, color: str, width: int) -> str:\n",
    "    \"\"\"Sets the stroke, color, and width of an element.\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "    if id not in elements:\n",
    "        raise ValueError(f\"Unknown id: {id}\")\n",
    "    elements[id].attrs[\"stroke\"] = color\n",
    "    elements[id].attrs[\"stroke-width\"] = width\n",
    "    return f\"Stroke of '{id}' set to {color}, width {width}\"\n",
    "\n",
    "@mcp.tool()\n",
    "def list_ids() -> List[str]:\n",
    "    \"\"\"Returns all element IDs currently in the SVG document.\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "    return list(elements.keys())\n",
    "\n",
    "@mcp.tool()\n",
    "def inspect_progress() -> str:\n",
    "    \"\"\"Inspects the current progress of the SVG drawing.\"\"\"\n",
    "    return \"Image will be shown to you in the next message.\"\n",
    "\n",
    "@mcp.tool()\n",
    "def export_svg() -> str:\n",
    "    \"\"\"Exports and returns the complete SVG document as an XML string.\"\"\"\n",
    "    canvas, elements = get_canvas_and_elements()\n",
    "    if not canvas:\n",
    "        raise ValueError(\"Canvas not created — call create_svg first\")\n",
    "    root = ET.Element(\"svg\", {\n",
    "        \"xmlns\": \"http://www.w3.org/2000/svg\",\n",
    "        \"width\": str(canvas[\"width\"]),\n",
    "        \"height\": str(canvas[\"height\"]),\n",
    "        \"viewBox\": f\"0 0 {canvas['width']} {canvas['height']}\"\n",
    "    })\n",
    "    for elem in elements.values():\n",
    "        ET.SubElement(root, elem.tag, {k: str(v) for k, v in elem.attrs.items()})\n",
    "    return ET.tostring(root, encoding=\"unicode\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "id": "6Xsu9Ywah1pd"
   },
   "source": [
    "# Code to evaluate production of valid unicorn and valid SVG\n",
    "\n",
    "*You should not need to modify this codeblock.*"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "id": "WnHaZ1NJL1m8"
   },
   "outputs": [],
   "source": [
    "REQUIRED_IDS = {\n",
    "    \"body\", \"head\", \"horn\", \"eye\",\n",
    "    \"leg1\", \"leg2\", \"leg3\", \"leg4\", \"tail\"\n",
    "}\n",
    "\n",
    "EXPECTED_TAGS = {\n",
    "    \"body\": \"ellipse\",\n",
    "    \"head\": \"circle\",\n",
    "    \"horn\": \"polygon\",\n",
    "    \"eye\": \"circle\",\n",
    "    \"leg1\": \"rect\",\n",
    "    \"leg2\": \"rect\",\n",
    "    \"leg3\": \"rect\",\n",
    "    \"leg4\": \"rect\",\n",
    "    \"tail\": \"polyline\",\n",
    "}\n",
    "\n",
    "def _strip_ns(tag: str) -> str:\n",
    "    return tag.split(\"}\")[-1]\n",
    "\n",
    "def _parse_points(points_str: str) -> List[Tuple[int, int]]:\n",
    "    pts = []\n",
    "    for item in points_str.strip().split():\n",
    "        x, y = item.split(\",\")\n",
    "        pts.append((int(float(x)), int(float(y))))\n",
    "    return pts\n",
    "\n",
    "def validate_unicornness(svg_string: str) -> Dict[str, Any]:\n",
    "    \"\"\"Automatic checks that the svg is well-formated and unicorn-like.\"\"\"\n",
    "\n",
    "    # Each check is binary (1 = pass, 0 = fail)\n",
    "    checks = {\n",
    "        \"canvas_size\": False,            # Correct canvas dimensions\n",
    "        \"required_ids_present\": False,   # All required elements present\n",
    "        \"element_types\": False,          # Each element has correct SVG tag\n",
    "        \"horn_above_head\": False,        # Horn is above head\n",
    "        \"legs_below_body_center\": False, # Legs are below body\n",
    "        \"eye_inside_head\": False,        # Eye lies inside head circle\n",
    "        \"all_colored\": False,            # Every element has fill or stroke\n",
    "        \"colorful\": False,               # At least three fill colors present\n",
    "        \"no_extra_ids\": False,           # No extra/unexpected elements\n",
    "        \"valid_xml\": False,              # SVG is valid XML\n",
    "    }\n",
    "\n",
    "    # Try parsing SVG string\n",
    "    try:\n",
    "        root = ET.fromstring(svg_string)\n",
    "        checks[\"valid_xml\"] = True\n",
    "    except Exception:\n",
    "        # If parsing fails, return early (all other checks fail)\n",
    "        return {\"checks\": checks, \"total\": sum(checks.values())}\n",
    "\n",
    "    # Check canvas size\n",
    "    width = root.attrib.get(\"width\")\n",
    "    height = root.attrib.get(\"height\")\n",
    "    if width == \"400\" and height == \"300\":\n",
    "        checks[\"canvas_size\"] = True\n",
    "\n",
    "    # Extract all elements with an \"id\"\n",
    "    elems = {}\n",
    "    for child in root:\n",
    "        elem_id = child.attrib.get(\"id\")\n",
    "        if elem_id:\n",
    "            elems[elem_id] = child\n",
    "\n",
    "    ids = set(elems.keys())\n",
    "\n",
    "    # Check if all required IDs are present\n",
    "    if REQUIRED_IDS.issubset(ids):\n",
    "        checks[\"required_ids_present\"] = True\n",
    "\n",
    "    # Check that there are no unexpected extra IDs\n",
    "    if ids.issubset(REQUIRED_IDS):\n",
    "        checks[\"no_extra_ids\"] = True\n",
    "\n",
    "    # Check that each required element exists and has the shape we'd expect\n",
    "    type_ok = True\n",
    "    for elem_id, tag in EXPECTED_TAGS.items():\n",
    "        if elem_id not in elems or _strip_ns(elems[elem_id].tag) != tag:\n",
    "            type_ok = False\n",
    "            break\n",
    "    if type_ok:\n",
    "        checks[\"element_types\"] = True\n",
    "\n",
    "    # Check that every element has either fill or stroke defined\n",
    "    colored_ok = True\n",
    "    for elem_id in REQUIRED_IDS:\n",
    "        if elem_id not in elems:\n",
    "            colored_ok = False\n",
    "            break\n",
    "        attrs = elems[elem_id].attrib\n",
    "        has_color = (\"fill\" in attrs) or (\"stroke\" in attrs)\n",
    "        if not has_color:\n",
    "            colored_ok = False\n",
    "            break\n",
    "    if colored_ok:\n",
    "        checks[\"all_colored\"] = True\n",
    "\n",
    "    # Check that the unicorn is colorful (has at least 3 distinct colors)\n",
    "    all_colors = set()\n",
    "    for elem_id, elem in elems.items():\n",
    "        attrs = elem.attrib\n",
    "        if \"fill\" in attrs and attrs.get(\"fill\") != \"none\":\n",
    "            all_colors.add(attrs.get(\"fill\").lower())\n",
    "    if len(all_colors) >= 3:\n",
    "        checks[\"colorful\"] = True\n",
    "\n",
    "    # Check that horn is above the head\n",
    "    if \"head\" in elems and \"horn\" in elems:\n",
    "        try:\n",
    "            head_cy = int(float(elems[\"head\"].attrib[\"cy\"]))\n",
    "            horn_pts = _parse_points(elems[\"horn\"].attrib[\"points\"])\n",
    "            if all(y < head_cy for _, y in horn_pts):\n",
    "                checks[\"horn_above_head\"] = True\n",
    "        except Exception:\n",
    "            pass\n",
    "\n",
    "    # Check that all legs are below the body center\n",
    "    if \"body\" in elems and all(f\"leg{i}\" in elems for i in range(1, 5)):\n",
    "        try:\n",
    "            body_cy = int(float(elems[\"body\"].attrib[\"cy\"]))\n",
    "            ok = True\n",
    "            for i in range(1, 5):\n",
    "                leg_y = int(float(elems[f\"leg{i}\"].attrib[\"y\"]))\n",
    "                if leg_y <= body_cy:\n",
    "                    ok = False\n",
    "                    break\n",
    "            if ok:\n",
    "                checks[\"legs_below_body_center\"] = True\n",
    "        except Exception:\n",
    "            pass\n",
    "\n",
    "    # Check that the eye lies inside the head circle\n",
    "    if \"head\" in elems and \"eye\" in elems:\n",
    "        try:\n",
    "            head_cx = int(float(elems[\"head\"].attrib[\"cx\"]))\n",
    "            head_cy = int(float(elems[\"head\"].attrib[\"cy\"]))\n",
    "            head_r = int(float(elems[\"head\"].attrib[\"r\"]))\n",
    "            eye_cx = int(float(elems[\"eye\"].attrib[\"cx\"]))\n",
    "            eye_cy = int(float(elems[\"eye\"].attrib[\"cy\"]))\n",
    "\n",
    "            # Euclidean distance between eye center and head center\n",
    "            dist = math.sqrt((eye_cx - head_cx) ** 2 + (eye_cy - head_cy) ** 2)\n",
    "\n",
    "            if dist < head_r:\n",
    "                checks[\"eye_inside_head\"] = True\n",
    "        except Exception:\n",
    "            pass\n",
    "\n",
    "    # Return individual check results + total score\n",
    "    return {\"checks\": checks, \"total\": sum(checks.values())}\n",
    "\n",
    "def summarize_eval_results(rows):\n",
    "    \"\"\"Given a list of eval checks, return a summary across them.\"\"\"\n",
    "    scores = [r[\"unicornness_score\"] for r in rows]\n",
    "    score_breakdowns = [r[\"score_breakdown\"] for r in rows]\n",
    "\n",
    "    avg_score = sum(scores) / len(rows)\n",
    "    perfect_runs = sum(score == 9 for score in scores)\n",
    "\n",
    "    # Count how many times each type of failure occurs.\n",
    "    fail_counts = {}\n",
    "    for key in score_breakdowns[0].keys():\n",
    "        fail_counts[key] = sum(1 - r[key] for r in score_breakdowns)\n",
    "\n",
    "    return {\n",
    "        \"avg_score\": avg_score,\n",
    "        \"perfect_runs\": perfect_runs,\n",
    "        \"fail_counts\": fail_counts\n",
    "    }"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "id": "sL85e9OArpfg"
   },
   "source": [
    "# Question 2.1"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "id": "Gik5Z-Hb3x0p"
   },
   "source": [
    "## Baseline agent\n",
    "\n",
    "The baseline agent doesn't use any language model. It draws a unicorn in a deterministic order, starting with the body, then the head, then the horn, then the legs, then the tail---each as a separate randomly-produced shape."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "id": "7NxJuBYkrrHm"
   },
   "outputs": [],
   "source": [
    "class BaselineAgent:\n",
    "    def __init__(self, seed: int = 0):\n",
    "        self.rng = random.Random(seed)\n",
    "        self.tool_log = list()\n",
    "\n",
    "    def _random_color(self):\n",
    "      return self.rng.choice([\"pink\", \"magenta\", \"purple\"])\n",
    "\n",
    "    def run(self):\n",
    "        create_svg(400, 300)\n",
    "\n",
    "        body_cx = 190 + self.rng.randint(-20, 20)\n",
    "        body_cy = 170 + self.rng.randint(-15, 15)\n",
    "        add_ellipse(body_cx, body_cy, 70, 40, \"body\")\n",
    "\n",
    "        head_cx = 280 + self.rng.randint(-25, 25)\n",
    "        head_cy = 125 + self.rng.randint(-25, 25)\n",
    "        add_circle(head_cx, head_cy, 30, \"head\")\n",
    "\n",
    "        horn_pts = [\n",
    "            [head_cx - 5, head_cy - 20],\n",
    "            [head_cx + 5, head_cy - 55 + self.rng.randint(-10, 20)],\n",
    "            [head_cx + 12, head_cy - 18]\n",
    "        ]\n",
    "        add_triangle(horn_pts, \"horn\")\n",
    "\n",
    "        add_circle(head_cx + 10, head_cy - 3, 4, \"eye\")\n",
    "\n",
    "        for i, x in enumerate([150, 180, 210, 240], start=1):\n",
    "            y = body_cy + self.rng.randint(-10, 40)\n",
    "            add_rect(x, y, 12, 55, f\"leg{i}\")\n",
    "\n",
    "        tail_pts = [\n",
    "            [120, body_cy - 10],\n",
    "            [90, body_cy - 25],\n",
    "            [75, body_cy]\n",
    "        ]\n",
    "        add_polyline(tail_pts, \"tail\")\n",
    "\n",
    "        for elem_id in list_ids():\n",
    "            if self.rng.random() < 0.85:\n",
    "                set_fill(elem_id, self._random_color())\n",
    "            if self.rng.random() < 0.50:\n",
    "                set_stroke(elem_id, \"black\", 2)\n",
    "\n",
    "        return export_svg()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "colab": {
     "base_uri": "https://localhost:8080/",
     "height": 1000
    },
    "id": "EREhXuO2rt6P",
    "outputId": "8667c4e7-4ac2-4c10-db38-043aa3771092"
   },
   "outputs": [],
   "source": [
    "def run_baseline(seeds):\n",
    "    \"\"\"Draw several unicorns using the baseline agent.\"\"\"\n",
    "    outputs = []\n",
    "    evals = []\n",
    "\n",
    "    for seed in seeds:\n",
    "        agent = BaselineAgent(seed=seed)\n",
    "        svg = agent.run()\n",
    "\n",
    "        eval_result = validate_unicornness(svg)\n",
    "\n",
    "        outputs.append(svg)\n",
    "        evals.append({\n",
    "            \"unicornness_score\": eval_result[\"total\"],\n",
    "            \"score_breakdown\": eval_result[\"checks\"],\n",
    "        })\n",
    "\n",
    "    return outputs, evals\n",
    "\n",
    "random_seeds = range(20)\n",
    "outputs_baseline, evals_baseline = run_baseline(seeds=random_seeds)\n",
    "\n",
    "print(json.dumps(summarize_eval_results(evals_baseline), indent=2))\n",
    "\n",
    "for n, (seed, evals, unicorn) in enumerate(zip(random_seeds, evals_baseline, outputs_baseline)):\n",
    "  print(f\"Attempt {n} (seed={seed})\")\n",
    "  print(json.dumps(evals, indent=2))\n",
    "  render_in_colab(unicorn)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "id": "QwqrFwrAryQT"
   },
   "source": [
    "# Question 2.2\n",
    "Your goal is to implement a PlanningAgent which does is guaranteed to satisfy all the conditions, producing the same unicorn every time."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "id": "OIR-AWQQr2fE"
   },
   "outputs": [],
   "source": [
    "class PlanningAgent:\n",
    "    def __init__(self, seed: int = 0):\n",
    "        self.seed = seed\n",
    "\n",
    "    def run(self):\n",
    "        \"\"\"\n",
    "        TODO:Make a deterministic sequence of tool calls.\n",
    "\n",
    "        Requirements:\n",
    "        - Include all required IDs exactly once\n",
    "        - Satisfy all validator constraints by construction\n",
    "        - Ensure every required element has fill or stroke\n",
    "        - Use of randomness is optional (so long as other requirements still hold)\n",
    "        \"\"\"\n",
    "        create_svg(400, 300)\n",
    "\n",
    "        # TODO 1:\n",
    "        # Add a call to create the SVG canvas with size 400 x 300\n",
    "\n",
    "        # TODO 2:\n",
    "        # Choose fixed coordinates for body and head\n",
    "        # Example variables:\n",
    "        # body_cx, body_cy = ...\n",
    "        # head_cx, head_cy = ...\n",
    "        # head_r = ...\n",
    "\n",
    "        # TODO 3:\n",
    "        # Add the body as an ellipse with id=\"body\"\n",
    "\n",
    "        # TODO 4:\n",
    "        # Add the head as a circle with id=\"head\"\n",
    "\n",
    "        # TODO 5:\n",
    "        # Add the horn as a triangle with id=\"horn\"\n",
    "        # Make sure all horn y-coordinates are strictly smaller than head_cy\n",
    "\n",
    "        # TODO 6:\n",
    "        # Add the eye as a circle with id=\"eye\"\n",
    "        # Make sure the eye center lies strictly inside the head circle\n",
    "\n",
    "        # TODO 7:\n",
    "        # Add four legs with ids leg1, leg2, leg3, leg4\n",
    "        # Make sure each leg y-coordinate is strictly greater than body_cy\n",
    "\n",
    "        # TODO 8:\n",
    "        # Add the tail as a polyline with id=\"tail\"\n",
    "\n",
    "        # TODO 9:\n",
    "        # Add styling calls so that every required element has fill or stroke\n",
    "        # You may assign both fill and stroke to every element\n",
    "\n",
    "        return export_svg()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "colab": {
     "base_uri": "https://localhost:8080/",
     "height": 910
    },
    "id": "oG2Y4Y0jr4dH",
    "outputId": "4b71eb5e-3462-4dbe-b11a-6d190be9fbbc"
   },
   "outputs": [],
   "source": [
    "def run_planning_agent(seeds):\n",
    "    \"\"\"Draw several unicorns using the planning agent.\"\"\"\n",
    "    outputs = []\n",
    "    evals = []\n",
    "\n",
    "    for seed in seeds:\n",
    "        agent = PlanningAgent(seed=seed)\n",
    "        svg = agent.run()\n",
    "\n",
    "        eval_result = validate_unicornness(svg)\n",
    "\n",
    "        outputs.append(svg)\n",
    "        evals.append({\n",
    "            \"unicornness_score\": eval_result[\"total\"],\n",
    "            \"score_breakdown\": eval_result[\"checks\"],\n",
    "        })\n",
    "\n",
    "    return outputs, evals\n",
    "\n",
    "random_seeds = range(1)\n",
    "outputs_planning_agent, evals_planning_agent = run_planning_agent(seeds=random_seeds)\n",
    "\n",
    "print(json.dumps(summarize_eval_results(evals_baseline), indent=2))\n",
    "\n",
    "for n, (seed, evals, unicorn) in enumerate(zip(random_seeds, evals_planning_agent, outputs_planning_agent)):\n",
    "  print(f\"Attempt {n} (seed={seed})\")\n",
    "  print(json.dumps(evals, indent=2))\n",
    "  render_in_colab(unicorn)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {
    "id": "mpRgxzXxr81q"
   },
   "source": [
    "# Question 2.3: Unicorn drawing with GPT-5"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "id": "xmLbCXvFfK0P"
   },
   "outputs": [],
   "source": [
    "async def get_tools_for_openai():\n",
    "    \"\"\"Returns a list of functions which are available for model to call.\"\"\"\n",
    "    async with Client(mcp) as client:\n",
    "        mcp_tools = await client.list_tools()\n",
    "    return [\n",
    "        {\n",
    "            \"type\": \"function\",\n",
    "            \"function\": {\n",
    "                \"name\": t.name,\n",
    "                \"description\": t.description,\n",
    "                \"parameters\": t.inputSchema,\n",
    "            }\n",
    "        }\n",
    "        for t in mcp_tools\n",
    "        if t.name not in [\"create_svg\", \"export_svg\"]\n",
    "    ]\n",
    "\n",
    "class OpenAIToolAgent:\n",
    "    def __init__(self, prompt:str):\n",
    "        self.prompt=prompt\n",
    "        self.width = 400\n",
    "        self.height = 300\n",
    "\n",
    "    async def run(self):\n",
    "        # Make sure you have your OpenAI key as a \"secret\" in Colab for this to work.\n",
    "        api_key = userdata.get('OPENAI_API_KEY')\n",
    "        openai_client = OpenAI(api_key=api_key)\n",
    "        tools = await get_tools_for_openai()\n",
    "\n",
    "        messages = [{\"role\": \"user\", \"content\": self.prompt}]\n",
    "\n",
    "        async with Client(mcp) as client:\n",
    "            # Instructor's note: we found that OpenAI likes to get stuck in a\n",
    "            # loop of calling `create_svg` repeatedly if `create_svg` is\n",
    "            # included in the MCP. Therefore, we call this manually before turning\n",
    "            # things over to the model and don't reveal ths function to the model.\n",
    "            await client.call_tool(\"create_svg\", {\"width\": self.width, \"height\": self.height})\n",
    "\n",
    "            temp = 0\n",
    "            while True:\n",
    "                response = openai_client.chat.completions.create(\n",
    "                    # You may choose to swap to a bigger/newer model if you'd like.\n",
    "                    model=\"gpt-5-mini\",\n",
    "                    tools=tools,\n",
    "                    messages=messages,\n",
    "                )\n",
    "\n",
    "                choice = response.choices[0]\n",
    "                messages.append(choice.message)  # Append the raw message object (OpenAI handles serialization)\n",
    "\n",
    "                # If done, break\n",
    "                if choice.finish_reason == \"stop\":\n",
    "                    print(f\"Model decides to finish.\")\n",
    "                    break\n",
    "\n",
    "                # Process tool calls\n",
    "                tool_calls = choice.message.tool_calls or []\n",
    "                for tool_call in tool_calls:\n",
    "                    name = tool_call.function.name\n",
    "                    args = json.loads(tool_call.function.arguments)\n",
    "\n",
    "                    print(f\"→ {name}({args})\")\n",
    "                    try:\n",
    "                        result = await client.call_tool(name, args)\n",
    "                        result_text = result[0].text if result else \"ok\"\n",
    "                    except Exception as e:\n",
    "                        result_text = f\"Error: {e}\"\n",
    "\n",
    "                    # Close out the tool call with plain text\n",
    "                    messages.append({\n",
    "                        \"role\": \"tool\",\n",
    "                        \"tool_call_id\": tool_call.id,\n",
    "                        \"content\": result_text,\n",
    "                    })\n",
    "\n",
    "                    if name == \"inspect_progress\":\n",
    "                        # Render directly from state\n",
    "                        client_response = await client.call_tool(\"export_svg\")\n",
    "                        svg_string = client_response.content[0].text\n",
    "                        png_bytes = cairosvg.svg2png(bytestring=svg_string.encode())\n",
    "                        b64 = base64.b64encode(png_bytes).decode(\"utf-8\")\n",
    "                        messages.append({\n",
    "                            \"role\": \"user\",\n",
    "                            \"content\": [\n",
    "                                {\n",
    "                                    \"type\": \"image_url\",\n",
    "                                    \"image_url\": {\n",
    "                                        \"url\": f\"data:image/png;base64,{b64}\"\n",
    "                                    }\n",
    "                                }\n",
    "                            ]\n",
    "                        })\n",
    "                        # Also display it in the notebook so you can follow along\n",
    "                        render_in_colab(svg_string)\n",
    "\n",
    "                if not tool_calls:\n",
    "                    break  # Safe exit if no tool calls and not \"stop\"\n",
    "\n",
    "            final_svg = await client.call_tool(\"export_svg\")\n",
    "            final_svg = final_svg.content[0].text\n",
    "        return messages, final_svg"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "colab": {
     "base_uri": "https://localhost:8080/",
     "height": 1000
    },
    "id": "rb5rGj8POlXg",
    "outputId": "b876fb82-1c86-4a05-e3ac-932c6af411cb"
   },
   "outputs": [],
   "source": [
    "prompt = f\"\"\"You have been provided an 400x300 canvas. Use the provided SVG tools to draw a beautiful, colorful unicorn onto it.\n",
    "    As minimum, your final unicorn should contain shapes with the IDs: {REQUIRED_IDS}.\n",
    "    At any point, you may call inspect_progress to visually inspect what the canvas currently looks like.\"\"\"\n",
    "\n",
    "agent = OpenAIToolAgent(prompt)\n",
    "messages, svg = await agent.run()\n",
    "print(\"Final unicorn:\")\n",
    "render_in_colab(svg)\n",
    "\n",
    "eval_result = validate_unicornness(svg)\n",
    "print(json.dumps(eval_result, indent=2))"
   ]
  }
 ],
 "metadata": {
  "colab": {
   "provenance": []
  },
  "kernelspec": {
   "display_name": "Python 3",
   "name": "python3"
  },
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 0
}
