{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Day 3: Guided Lab — Building a RAG Pipeline\n",
    "\n",
    "## Learning Objectives\n",
    "\n",
    "1. Create text embeddings using the Gemini Embeddings API\n",
    "2. Implement cosine similarity and nearest-neighbor search\n",
    "3. Split documents into semantically meaningful chunks\n",
    "4. Build a complete RAG pipeline (retrieve → augment → generate)\n",
    "5. Evaluate RAG outputs using the RAG triad framework\n",
    "\n",
    "## Prerequisites\n",
    "\n",
    "- [x] Completed Day 2 labs\n",
    "- [x] Understanding of embeddings from input session\n",
    "- [x] Google Colab with `GEMINI_API_KEY` configured"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Part 0: Setup and Infrastructure"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "!pip install -q -U google-genai"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# ── Imports ──────────────────────────────────────────────\n",
    "import os, time, json\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "from datetime import datetime, timezone\n",
    "from typing import List, Optional, Literal\n",
    "from pydantic import BaseModel, Field\n",
    "\n",
    "from google import genai\n",
    "from google.genai import types\n",
    "\n",
    "# ── API Key ───────────────────────────────────────────────\n",
    "try:\n",
    "    from google.colab import userdata\n",
    "    API_KEY = userdata.get(\"GEMINI_API_KEY\")\n",
    "except Exception:\n",
    "    API_KEY = None\n",
    "\n",
    "if not API_KEY:\n",
    "    import getpass\n",
    "    API_KEY = getpass.getpass(\"Enter your Gemini API key: \")\n",
    "\n",
    "client = genai.Client(api_key=API_KEY)\n",
    "\n",
    "MODEL_ID = \"gemini-2.5-flash-lite\"\n",
    "EMBEDDING_MODEL = \"gemini-embedding-001\"\n",
    "\n",
    "print(f\"API key loaded: {'yes' if API_KEY else 'no'}\")\n",
    "print(f\"Generation model: {MODEL_ID}\")\n",
    "print(f\"Embedding model:  {EMBEDDING_MODEL}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# ── Logging Infrastructure ────────────────────────────────\n",
    "PROMPT_LOG = []\n",
    "\n",
    "def _now():\n",
    "    return datetime.now(timezone.utc).isoformat().replace('+00:00', 'Z')\n",
    "\n",
    "def generate(prompt, temperature=0.7, max_tokens=1000, log=True, label=None):\n",
    "    \"\"\"Generate free-form text. Returns raw string.\"\"\"\n",
    "    t0 = time.time()\n",
    "    response = client.models.generate_content(\n",
    "        model=MODEL_ID,\n",
    "        contents=prompt,\n",
    "        config={\"temperature\": temperature, \"max_output_tokens\": max_tokens},\n",
    "    )\n",
    "    latency = time.time() - t0\n",
    "    text = response.text\n",
    "    if log:\n",
    "        PROMPT_LOG.append({\n",
    "            \"timestamp\": _now(), \"label\": label or \"generate\",\n",
    "            \"type\": \"free_form\",\n",
    "            \"prompt\": prompt[:300] + \"...\" if len(prompt) > 300 else prompt,\n",
    "            \"prompt_length\": len(prompt),\n",
    "            \"temperature\": temperature,\n",
    "            \"response\": text[:300] + \"...\" if len(text) > 300 else text,\n",
    "            \"response_length\": len(text),\n",
    "            \"latency_s\": round(latency, 2),\n",
    "        })\n",
    "    return text\n",
    "\n",
    "def generate_structured(prompt, schema_model, temperature=0.2, log=True, label=None):\n",
    "    \"\"\"Generate structured JSON output. Returns Pydantic model instance.\"\"\"\n",
    "    t0 = time.time()\n",
    "    response = client.models.generate_content(\n",
    "        model=MODEL_ID,\n",
    "        contents=prompt,\n",
    "        config={\n",
    "            \"temperature\": temperature,\n",
    "            \"response_mime_type\": \"application/json\",\n",
    "            \"response_schema\": schema_model,\n",
    "        },\n",
    "    )\n",
    "    latency = time.time() - t0\n",
    "    result = schema_model.model_validate_json(response.text)\n",
    "    if log:\n",
    "        PROMPT_LOG.append({\n",
    "            \"timestamp\": _now(), \"label\": label or \"structured\",\n",
    "            \"type\": \"structured\",\n",
    "            \"schema\": schema_model.__name__,\n",
    "            \"prompt\": prompt[:300] + \"...\" if len(prompt) > 300 else prompt,\n",
    "            \"prompt_length\": len(prompt),\n",
    "            \"temperature\": temperature,\n",
    "            \"response\": response.text[:300],\n",
    "            \"response_length\": len(response.text),\n",
    "            \"latency_s\": round(latency, 2),\n",
    "        })\n",
    "    return result\n",
    "\n",
    "def show_log():\n",
    "    \"\"\"Display prompt log as DataFrame.\"\"\"\n",
    "    if not PROMPT_LOG:\n",
    "        print(\"No API calls logged yet.\")\n",
    "        return None\n",
    "    return pd.DataFrame(PROMPT_LOG)\n",
    "\n",
    "print(\"Logging infrastructure ready.\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# ── RAG Infrastructure ────────────────────────────────────\n",
    "\n",
    "def embed_texts(texts, task_type=\"RETRIEVAL_DOCUMENT\"):\n",
    "    \"\"\"\n",
    "    Embed one or more texts using Gemini Embeddings API.\n",
    "    \n",
    "    Args:\n",
    "        texts: A string or list of strings to embed.\n",
    "        task_type: 'RETRIEVAL_DOCUMENT', 'RETRIEVAL_QUERY', or 'SEMANTIC_SIMILARITY'.\n",
    "    \n",
    "    Returns:\n",
    "        List of numpy arrays (one embedding per input text).\n",
    "    \"\"\"\n",
    "    if isinstance(texts, str):\n",
    "        texts = [texts]\n",
    "    response = client.models.embed_content(\n",
    "        model=EMBEDDING_MODEL,\n",
    "        contents=texts,\n",
    "        config=types.EmbedContentConfig(task_type=task_type),\n",
    "    )\n",
    "    return [np.array(e.values) for e in response.embeddings]\n",
    "\n",
    "\n",
    "def cosine_similarity(a, b):\n",
    "    \"\"\"Compute cosine similarity between two vectors.\"\"\"\n",
    "    return float(np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b)))\n",
    "\n",
    "\n",
    "def search(query, doc_embeddings, documents, top_k=3):\n",
    "    \"\"\"\n",
    "    Find top-k most similar documents to a query.\n",
    "    \n",
    "    Returns: list of (index, score, document_text) tuples, sorted by score descending.\n",
    "    \"\"\"\n",
    "    query_emb = embed_texts(query, task_type=\"RETRIEVAL_QUERY\")[0]\n",
    "    scores = []\n",
    "    for i, doc_emb in enumerate(doc_embeddings):\n",
    "        sim = cosine_similarity(query_emb, doc_emb)\n",
    "        scores.append((i, sim, documents[i]))\n",
    "    scores.sort(key=lambda x: x[1], reverse=True)\n",
    "    return scores[:top_k]\n",
    "\n",
    "\n",
    "def rag_query(question, documents, doc_embeddings, top_k=3, system_prompt=None):\n",
    "    \"\"\"\n",
    "    Complete RAG pipeline: retrieve → augment → generate.\n",
    "    \n",
    "    Returns: (answer_text, retrieved_results)\n",
    "    \"\"\"\n",
    "    # Retrieve\n",
    "    results = search(question, doc_embeddings, documents, top_k)\n",
    "    context = \"\\n\\n\".join(\n",
    "        f\"[Chunk {i+1}]\\n{doc}\" for i, (_, _, doc) in enumerate(results)\n",
    "    )\n",
    "    \n",
    "    # Augment\n",
    "    if system_prompt is None:\n",
    "        system_prompt = (\n",
    "            \"You are a helpful assistant that answers questions based on \"\n",
    "            \"the provided context. If the answer is not in the context, \"\n",
    "            'say \"I don\\'t have enough information to answer this.\" '\n",
    "            \"Do not make up information.\"\n",
    "        )\n",
    "    \n",
    "    prompt = f\"\"\"{system_prompt}\n",
    "\n",
    "<context>\n",
    "{context}\n",
    "</context>\n",
    "\n",
    "Question: {question}\"\"\"\n",
    "    \n",
    "    # Generate\n",
    "    answer = generate(prompt, temperature=0.3, label=f\"rag_{question[:30]}\")\n",
    "    return answer, results\n",
    "\n",
    "print(\"RAG infrastructure ready: embed_texts, cosine_similarity, search, rag_query\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Part 1: Text Embeddings\n",
    "\n",
    "Embeddings convert text into numerical vectors where **similar meanings are close together**.\n",
    "We'll use Google's `gemini-embedding-001` model which produces 3072-dimensional vectors."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 1.1: Embed Documents\n",
    "\n",
    "Let's embed a set of company documents and examine what the embedding vectors look like."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Our sample knowledge base (6 company documents)\n",
    "documents = [\n",
    "    \"The company was founded in 2015 by Maria Chen and David Park. \"\n",
    "    \"It started as a two-person startup in a garage in Munich.\",\n",
    "\n",
    "    \"Our remote work policy allows employees to work from home up to \"\n",
    "    \"3 days per week. All team meetings on Tuesdays and Thursdays are mandatory in-person.\",\n",
    "\n",
    "    \"Q3 2025 revenue was EUR 4.2 million, a 15% increase over Q3 2024. \"\n",
    "    \"Growth was primarily driven by the enterprise segment.\",\n",
    "\n",
    "    \"The company offers 30 days of annual leave plus public holidays. \"\n",
    "    \"Unused leave can be carried over for up to 6 months.\",\n",
    "\n",
    "    \"Our main product, DataSync Pro, integrates with SAP, Salesforce, \"\n",
    "    \"and Microsoft 365. The API supports REST and GraphQL.\",\n",
    "\n",
    "    \"The engineering team follows a two-week sprint cycle with planning \"\n",
    "    \"on Mondays and retrospectives on Fridays.\",\n",
    "]\n",
    "\n",
    "# Embed all documents\n",
    "doc_embeddings = embed_texts(documents, task_type=\"RETRIEVAL_DOCUMENT\")\n",
    "\n",
    "print(f\"Number of documents: {len(doc_embeddings)}\")\n",
    "print(f\"Embedding dimension: {len(doc_embeddings[0])}\")\n",
    "print(f\"\\nFirst 5 values of document 0: {doc_embeddings[0][:5]}\")\n",
    "print(f\"Vector norm (length):          {np.linalg.norm(doc_embeddings[0]):.4f}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 1.2: Compare Task Types\n",
    "\n",
    "The Gemini Embeddings API supports different **task types** that optimise the embedding for its intended use."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "query_text = \"Can I work from home?\"\n",
    "\n",
    "# Embed the same text with different task types\n",
    "emb_doc   = embed_texts(query_text, task_type=\"RETRIEVAL_DOCUMENT\")[0]\n",
    "emb_query = embed_texts(query_text, task_type=\"RETRIEVAL_QUERY\")[0]\n",
    "emb_sim   = embed_texts(query_text, task_type=\"SEMANTIC_SIMILARITY\")[0]\n",
    "\n",
    "print(f\"Same text, different task types:\")\n",
    "print(f\"  DOC  vs QUERY:      cosine = {cosine_similarity(emb_doc, emb_query):.4f}\")\n",
    "print(f\"  DOC  vs SIMILARITY: cosine = {cosine_similarity(emb_doc, emb_sim):.4f}\")\n",
    "print(f\"  QUERY vs SIMILARITY: cosine = {cosine_similarity(emb_query, emb_sim):.4f}\")\n",
    "print()\n",
    "print(\"Notice: The same text produces DIFFERENT vectors depending on task type!\")\n",
    "print(\"Use RETRIEVAL_DOCUMENT for documents, RETRIEVAL_QUERY for queries.\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "> **💡 Discussion:** Each embedding is a 3072-dimensional vector. You can think of each dimension as capturing one aspect of meaning. Similar texts will have similar patterns across all 3072 dimensions, which is why cosine similarity works — it measures how much two vectors \"point in the same direction\" in this high-dimensional space."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Part 2: Similarity Search\n",
    "\n",
    "Now that we have embeddings, we can find documents by **meaning** instead of keywords."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 2.1: Pairwise Similarity\n",
    "\n",
    "Let's see how similar our 6 documents are to each other."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Compute pairwise cosine similarity\n",
    "n = len(doc_embeddings)\n",
    "sim_matrix = np.zeros((n, n))\n",
    "for i in range(n):\n",
    "    for j in range(n):\n",
    "        sim_matrix[i][j] = cosine_similarity(doc_embeddings[i], doc_embeddings[j])\n",
    "\n",
    "# Display as DataFrame\n",
    "labels = [f\"Doc {i}\" for i in range(n)]\n",
    "sim_df = pd.DataFrame(sim_matrix, index=labels, columns=labels).round(3)\n",
    "print(\"Pairwise cosine similarity between documents:\")\n",
    "print(sim_df.to_string())\n",
    "print()\n",
    "\n",
    "# Find most similar pair (excluding self-comparisons)\n",
    "np.fill_diagonal(sim_matrix, 0)\n",
    "max_idx = np.unravel_index(sim_matrix.argmax(), sim_matrix.shape)\n",
    "print(f\"Most similar pair: Doc {max_idx[0]} and Doc {max_idx[1]} \"\n",
    "      f\"(cosine = {sim_matrix[max_idx]:.3f})\")\n",
    "print(f\"  Doc {max_idx[0]}: {documents[max_idx[0]][:60]}...\")\n",
    "print(f\"  Doc {max_idx[1]}: {documents[max_idx[1]][:60]}...\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 2.2: Semantic Search\n",
    "\n",
    "Search our documents by meaning using the `search()` function."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "queries = [\n",
    "    \"Can I work from home?\",\n",
    "    \"How much money did we make?\",\n",
    "    \"Who started the company?\",\n",
    "    \"What tools does the product connect to?\",\n",
    "]\n",
    "\n",
    "for q in queries:\n",
    "    results = search(q, doc_embeddings, documents, top_k=2)\n",
    "    print(f\"\\nQuery: \\\"{q}\\\"\")\n",
    "    for idx, score, doc in results:\n",
    "        print(f\"  [{score:.3f}] Doc {idx}: {doc[:70]}...\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Part 3: Document Chunking\n",
    "\n",
    "Real documents are long. We need to **split them into smaller chunks** before embedding.\n",
    "Let's explore different chunking strategies."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 3.1: Fixed-Size Chunking"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# A longer document to chunk\n",
    "long_document = \"\"\"Remote Work Policy (Version 4.0, Updated January 2025)\n",
    "\n",
    "Section 1: General Guidelines\n",
    "Our company supports flexible work arrangements to help employees balance \n",
    "productivity and well-being. Employees may work remotely up to 3 days per \n",
    "week, subject to manager approval. New employees in their first 90 days \n",
    "must work on-site full-time to complete onboarding.\n",
    "\n",
    "Section 2: Mandatory In-Office Days\n",
    "All team-wide meetings are held on Tuesdays and Thursdays. Attendance is \n",
    "mandatory and in-person. Department heads may designate additional in-office \n",
    "days for project-critical phases. Failure to attend mandatory meetings \n",
    "without prior approval will be noted in performance reviews.\n",
    "\n",
    "Section 3: Equipment and Expenses\n",
    "The company provides a one-time home office stipend of EUR 500 for ergonomic \n",
    "equipment (desk, chair, monitor). Internet costs are reimbursed up to \n",
    "EUR 30 per month upon submission of receipts. All company equipment \n",
    "must be returned upon termination.\n",
    "\n",
    "Section 4: Security Requirements\n",
    "Remote workers must use the company VPN at all times when accessing \n",
    "internal systems. Sensitive documents must not be printed at home. \n",
    "Screen locks must engage after 5 minutes of inactivity. Any security \n",
    "incidents must be reported to IT within 1 hour.\"\"\"\n",
    "\n",
    "# Fixed-size chunking (by word count)\n",
    "def chunk_fixed(text, chunk_size=50, overlap=10):\n",
    "    \"\"\"Split text into chunks of approximately chunk_size words with overlap.\"\"\"\n",
    "    words = text.split()\n",
    "    chunks = []\n",
    "    start = 0\n",
    "    while start < len(words):\n",
    "        end = start + chunk_size\n",
    "        chunk = \" \".join(words[start:end])\n",
    "        chunks.append(chunk)\n",
    "        start = end - overlap  # Overlap\n",
    "    return chunks\n",
    "\n",
    "fixed_chunks = chunk_fixed(long_document, chunk_size=50, overlap=10)\n",
    "print(f\"Fixed-size chunking: {len(fixed_chunks)} chunks\\n\")\n",
    "for i, chunk in enumerate(fixed_chunks):\n",
    "    print(f\"Chunk {i} ({len(chunk.split())} words): {chunk[:80]}...\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 3.2: Sentence-Based Chunking"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Sentence-based chunking (respects sentence boundaries)\n",
    "def chunk_sentences(text, max_words=60, overlap_words=15):\n",
    "    \"\"\"Split text at sentence boundaries, grouping sentences up to max_words.\"\"\"\n",
    "    # Simple sentence split (handles common abbreviations poorly, but fine for demo)\n",
    "    sentences = [s.strip() for s in text.replace('\\n', ' ').split('. ') if s.strip()]\n",
    "    \n",
    "    chunks = []\n",
    "    current = []\n",
    "    current_len = 0\n",
    "    \n",
    "    for sent in sentences:\n",
    "        sent_words = len(sent.split())\n",
    "        if current_len + sent_words > max_words and current:\n",
    "            chunks.append(\". \".join(current) + \".\")\n",
    "            # Keep last sentence(s) as overlap\n",
    "            overlap_sents = []\n",
    "            overlap_len = 0\n",
    "            for s in reversed(current):\n",
    "                if overlap_len + len(s.split()) <= overlap_words:\n",
    "                    overlap_sents.insert(0, s)\n",
    "                    overlap_len += len(s.split())\n",
    "                else:\n",
    "                    break\n",
    "            current = overlap_sents\n",
    "            current_len = overlap_len\n",
    "        current.append(sent)\n",
    "        current_len += sent_words\n",
    "    \n",
    "    if current:\n",
    "        chunks.append(\". \".join(current) + \".\")\n",
    "    return chunks\n",
    "\n",
    "sent_chunks = chunk_sentences(long_document, max_words=60, overlap_words=15)\n",
    "print(f\"Sentence-based chunking: {len(sent_chunks)} chunks\\n\")\n",
    "for i, chunk in enumerate(sent_chunks):\n",
    "    print(f\"Chunk {i} ({len(chunk.split())} words):\")\n",
    "    print(f\"  {chunk[:120]}...\\n\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 3.3: Adding Chunk Metadata\n",
    "\n",
    "In production, you need to know **where each chunk came from** — which document, which section, when it was last updated. Let's attach metadata to our chunks."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Attach metadata to each chunk (production pattern)\n",
    "from dataclasses import dataclass\n",
    "\n",
    "@dataclass\n",
    "class Chunk:\n",
    "    \"\"\"A chunk with metadata for traceability.\"\"\"\n",
    "    chunk_id: str\n",
    "    doc_title: str\n",
    "    text: str\n",
    "\n",
    "# Create metadata-aware chunks from our policy document\n",
    "policy_title = \"Remote Work Policy v4.0\"\n",
    "metadata_chunks = []\n",
    "for i, text in enumerate(sent_chunks):\n",
    "    metadata_chunks.append(Chunk(\n",
    "        chunk_id=f\"REMOTE_WORK::C{i+1}\",\n",
    "        doc_title=policy_title,\n",
    "        text=text,\n",
    "    ))\n",
    "\n",
    "print(f\"Created {len(metadata_chunks)} chunks with metadata:\\n\")\n",
    "for c in metadata_chunks:\n",
    "    print(f\"  {c.chunk_id} ({c.doc_title})\")\n",
    "    print(f\"    {c.text[:80]}...\\n\")\n",
    "\n",
    "print(\"Now each chunk carries its identity — essential for citations and debugging.\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "> **💡 Discussion: Chunk Size Trade-offs**\n",
    ">\n",
    "> | | Small Chunks (~50 words) | Large Chunks (~150 words) |\n",
    "> |---|---|---|\n",
    "> | **Precision** | High — focused content | Lower — includes noise |\n",
    "> | **Context** | Risk missing surrounding info | Better coverage |\n",
    "> | **Number** | Many chunks to search | Fewer, more efficient |\n",
    ">\n",
    "> **Rule of thumb:** Start with 200–500 tokens (~50–125 words) with 10–20% overlap."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Part 4: The Complete RAG Pipeline\n",
    "\n",
    "Now let's put it all together: **retrieve** relevant chunks, **augment** the prompt, and **generate** a grounded answer."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 4.1: RAG with Grounding Rules"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Use the sentence-based chunks from our remote work policy\n",
    "policy_chunks = sent_chunks\n",
    "policy_embeddings = embed_texts(policy_chunks, task_type=\"RETRIEVAL_DOCUMENT\")\n",
    "\n",
    "# Define a strong grounding prompt\n",
    "SYSTEM_PROMPT = \"\"\"You are a company policy assistant.\n",
    "Answer questions based ONLY on the provided policy excerpts.\n",
    "Rules:\n",
    "- If the answer is not in the excerpts, say: \"I don't have that information in the provided policies.\"\n",
    "- Cite which chunk(s) support your answer using [Chunk N] notation.\n",
    "- Keep answers concise (under 80 words).\n",
    "- Do not invent policy details.\"\"\"\n",
    "\n",
    "# Ask a question\n",
    "question = \"How many days can I work from home?\"\n",
    "answer, retrieved = rag_query(\n",
    "    question, policy_chunks, policy_embeddings,\n",
    "    top_k=3, system_prompt=SYSTEM_PROMPT\n",
    ")\n",
    "\n",
    "print(f\"Question: {question}\\n\")\n",
    "print(\"Retrieved chunks:\")\n",
    "for idx, score, doc in retrieved:\n",
    "    print(f\"  [{score:.3f}] Chunk {idx}: {doc[:70]}...\")\n",
    "print(f\"\\nAnswer:\\n{answer}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 4.2: With vs. Without Context (Hallucination Demo)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "question = \"What is the home office stipend amount?\"\n",
    "\n",
    "# WITH RAG (grounded)\n",
    "rag_answer, _ = rag_query(\n",
    "    question, policy_chunks, policy_embeddings,\n",
    "    top_k=3, system_prompt=SYSTEM_PROMPT\n",
    ")\n",
    "\n",
    "# WITHOUT RAG (ungrounded — model must rely on \"memory\")\n",
    "no_rag_answer = generate(\n",
    "    f\"Answer this company policy question: {question}\",\n",
    "    temperature=0.3, label=\"no_rag_comparison\"\n",
    ")\n",
    "\n",
    "print(f\"Question: {question}\\n\")\n",
    "print(f\"WITH RAG (grounded):\\n{rag_answer}\\n\")\n",
    "print(f\"WITHOUT RAG (ungrounded):\\n{no_rag_answer}\\n\")\n",
    "print(\"Notice: Without context, the model either guesses or admits it doesn't know.\")\n",
    "print(\"With RAG, the answer is grounded in actual policy text.\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 4.3: Out-of-Scope Question (Refusal Test)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# This question has NO answer in our policy documents\n",
    "out_of_scope = \"What programming language is the backend written in?\"\n",
    "\n",
    "answer, retrieved = rag_query(\n",
    "    out_of_scope, policy_chunks, policy_embeddings,\n",
    "    top_k=3, system_prompt=SYSTEM_PROMPT\n",
    ")\n",
    "\n",
    "print(f\"Question: {out_of_scope}\\n\")\n",
    "print(f\"Answer:\\n{answer}\\n\")\n",
    "print(\"The system should refuse to answer (not hallucinate a programming language).\")"
   ]
  },
  {
   "cell_type": "markdown",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": "---\n## Part 5: RAG Evaluation\n\nMeasure RAG quality systematically using a golden test set, keyword matching, retrieval metrics, and LLM-as-judge scoring."
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Exercise 5.1: Define a Golden Test Set"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": "# Golden set: questions with expected answers, source chunks, and expected chunk indices\nGOLDEN_SET = [\n    {\n        \"id\": \"Q1\",\n        \"question\": \"How many days per week can I work remotely?\",\n        \"expected_keywords\": [\"3 days\", \"three days\"],\n        \"expected_chunks\": [0, 1],  # Chunk indices that should be retrieved\n        \"difficulty\": \"easy\",\n    },\n    {\n        \"id\": \"Q2\",\n        \"question\": \"Which days must I be in the office?\",\n        \"expected_keywords\": [\"Tuesday\", \"Thursday\"],\n        \"expected_chunks\": [1],\n        \"difficulty\": \"easy\",\n    },\n    {\n        \"id\": \"Q3\",\n        \"question\": \"How much is the home office equipment budget?\",\n        \"expected_keywords\": [\"500\", \"EUR 500\"],\n        \"expected_chunks\": [2, 3],\n        \"difficulty\": \"easy\",\n    },\n    {\n        \"id\": \"Q4\",\n        \"question\": \"Can new employees work from home immediately?\",\n        \"expected_keywords\": [\"90 days\", \"onboarding\", \"first 90\"],\n        \"expected_chunks\": [0, 1],\n        \"difficulty\": \"medium\",\n    },\n    {\n        \"id\": \"Q5\",\n        \"question\": \"What happens if I miss a mandatory meeting?\",\n        \"expected_keywords\": [\"performance review\"],\n        \"expected_chunks\": [1, 2],\n        \"difficulty\": \"medium\",\n    },\n    {\n        \"id\": \"Q6\",\n        \"question\": \"How quickly must security incidents be reported?\",\n        \"expected_keywords\": [\"1 hour\", \"one hour\"],\n        \"expected_chunks\": [3, 4],\n        \"difficulty\": \"medium\",\n    },\n    {\n        \"id\": \"Q7\",\n        \"question\": \"What is the company's revenue forecast for 2026?\",\n        \"expected_keywords\": [\"don't have\", \"not in\", \"no information\"],\n        \"expected_chunks\": [],  # No relevant chunk exists\n        \"difficulty\": \"refusal\",\n    },\n    {\n        \"id\": \"Q8\",\n        \"question\": \"What's the best restaurant near the office?\",\n        \"expected_keywords\": [\"don't have\", \"not in\", \"no information\"],\n        \"expected_chunks\": [],\n        \"difficulty\": \"refusal\",\n    },\n]\n\nprint(f\"Golden set: {len(GOLDEN_SET)} questions\")\nfor q in GOLDEN_SET:\n    chunks_str = ', '.join(str(c) for c in q['expected_chunks']) if q['expected_chunks'] else 'none'\n    print(f\"  [{q['difficulty']:7s}] {q['id']}: {q['question']}  (chunks: {chunks_str})\")"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "### Exercise 5.2: Keyword-Based Evaluation\n\nRun the RAG pipeline on all golden set questions and check whether the answers contain the expected keywords."
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": "# Run RAG on all golden set questions\nrag_results = []\n\nfor qa in GOLDEN_SET:\n    answer, retrieved = rag_query(\n        qa[\"question\"], policy_chunks, policy_embeddings,\n        top_k=3, system_prompt=SYSTEM_PROMPT\n    )\n    \n    # Check if answer contains expected keywords\n    answer_lower = answer.lower()\n    keyword_hit = any(\n        kw.lower() in answer_lower for kw in qa[\"expected_keywords\"]\n    )\n    \n    rag_results.append({\n        \"id\": qa[\"id\"],\n        \"question\": qa[\"question\"],\n        \"difficulty\": qa[\"difficulty\"],\n        \"answer\": answer,\n        \"keyword_match\": keyword_hit,\n        \"top_chunk_score\": retrieved[0][1],\n    })\n\n# Summary\nresults_df = pd.DataFrame(rag_results)\naccuracy = results_df[\"keyword_match\"].mean()\nprint(f\"\\nOverall keyword accuracy: {accuracy:.0%} ({results_df['keyword_match'].sum()}/{len(results_df)})\")\nprint()\nprint(results_df[[\"id\", \"difficulty\", \"keyword_match\", \"top_chunk_score\"]].to_string(index=False))"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "### Exercise 5.3: Precision@k and Recall@k\n\nBeyond keyword matching, we can measure **retrieval quality** directly using standard IR metrics:\n\n- **Precision@k** = What fraction of retrieved chunks were relevant?\n- **Recall@k** = What fraction of relevant chunks were retrieved?"
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": "# Precision@k and Recall@k\ndef precision_recall_at_k(retrieved_indices, expected_chunks, k):\n    \"\"\"Compute Precision@k and Recall@k for a single query.\"\"\"\n    retrieved_set = set(retrieved_indices[:k])\n    expected_set = set(expected_chunks)\n    if not expected_set:  # Refusal questions have no expected chunks\n        return None, None\n    hits = retrieved_set & expected_set\n    precision = len(hits) / k if k > 0 else 0\n    recall = len(hits) / len(expected_set) if expected_set else 0\n    return precision, recall\n\n# Compute for all non-refusal golden set questions\nretrieval_metrics = []\nfor qa in GOLDEN_SET:\n    if not qa[\"expected_chunks\"]:  # Skip refusal questions\n        continue\n    # Get retrieved chunk indices\n    results = search(qa[\"question\"], policy_embeddings, policy_chunks, top_k=3)\n    retrieved_indices = [idx for idx, _, _ in results]\n    \n    p, r = precision_recall_at_k(retrieved_indices, qa[\"expected_chunks\"], k=3)\n    retrieval_metrics.append({\n        \"id\": qa[\"id\"],\n        \"question\": qa[\"question\"][:40],\n        \"precision@3\": p,\n        \"recall@3\": r,\n        \"retrieved\": retrieved_indices,\n        \"expected\": qa[\"expected_chunks\"],\n    })\n\nmetrics_df = pd.DataFrame(retrieval_metrics)\nprint(\"Retrieval Metrics (Precision@k / Recall@k):\")\nprint(metrics_df[[\"id\", \"precision@3\", \"recall@3\"]].to_string(index=False))\nprint(f\"\\nAvg Precision@3: {metrics_df['precision@3'].mean():.2f}\")\nprint(f\"Avg Recall@3:    {metrics_df['recall@3'].mean():.2f}\")"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "### Exercise 5.4: LLM-as-Judge Scoring"
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Use the LLM to score answer quality\n",
    "class RAGScore(BaseModel):\n",
    "    \"\"\"RAG quality score for a single answer.\"\"\"\n",
    "    groundedness: int = Field(description=\"1-5: Is the answer supported by the context?\")\n",
    "    relevance: int = Field(description=\"1-5: Does the answer address the question?\")\n",
    "    explanation: str = Field(description=\"Brief explanation of the scores\")\n",
    "\n",
    "# Score a few answers\n",
    "scores = []\n",
    "for r in rag_results[:5]:  # Score first 5 to save API calls\n",
    "    eval_prompt = f\"\"\"Rate this RAG system answer.\n",
    "\n",
    "Question: {r['question']}\n",
    "Answer: {r['answer']}\n",
    "\n",
    "Score groundedness (1-5): Is the answer based on facts, not made up?\n",
    "Score relevance (1-5): Does the answer address the question?\"\"\"\n",
    "    \n",
    "    score = generate_structured(eval_prompt, RAGScore, label=f\"eval_{r['id']}\")\n",
    "    scores.append({\n",
    "        \"id\": r[\"id\"],\n",
    "        \"groundedness\": score.groundedness,\n",
    "        \"relevance\": score.relevance,\n",
    "        \"explanation\": score.explanation,\n",
    "    })\n",
    "\n",
    "scores_df = pd.DataFrame(scores)\n",
    "print(\"LLM-as-Judge Scores:\")\n",
    "print(scores_df.to_string(index=False))\n",
    "print(f\"\\nAvg Groundedness: {scores_df['groundedness'].mean():.1f}/5\")\n",
    "print(f\"Avg Relevance:    {scores_df['relevance'].mean():.1f}/5\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "---\n## Summary and Key Takeaways\n\n| Technique | What You Learned | Key Function |\n|-----------|-----------------|---------------|\n| **Text Embeddings** | Convert text to 3072-dim vectors | `embed_texts()` |\n| **Task Types** | RETRIEVAL_DOCUMENT vs RETRIEVAL_QUERY | `task_type` parameter |\n| **Cosine Similarity** | Measure semantic closeness (0–1) | `cosine_similarity()` |\n| **Document Chunking** | Split long docs for precise retrieval | `chunk_fixed()`, `chunk_sentences()` |\n| **Chunk Metadata** | Attach doc_id, title for traceability | `@dataclass Chunk` |\n| **RAG Pipeline** | Retrieve → Augment → Generate | `rag_query()` |\n| **Grounding Prompts** | Force model to use only context | System prompt rules |\n| **Golden Test Set** | Systematic evaluation of RAG quality | Expected keywords + LLM-as-judge |\n| **Precision@k / Recall@k** | Measure retrieval quality directly | Standard IR metrics |\n\n### Checklist\n\n- [x] Embedded documents and examined vector properties\n- [x] Compared RETRIEVAL_DOCUMENT vs RETRIEVAL_QUERY task types\n- [x] Computed pairwise similarity between documents\n- [x] Searched by meaning (semantic search)\n- [x] Chunked a longer document two ways\n- [x] Attached metadata to chunks for traceability\n- [x] Built a complete RAG pipeline with grounding prompt\n- [x] Tested with vs. without context (hallucination demo)\n- [x] Tested refusal on out-of-scope questions\n- [x] Evaluated with golden test set and keyword matching\n- [x] Computed Precision@k and Recall@k retrieval metrics\n- [x] Scored answers with LLM-as-judge"
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# ── Export Experiment Log ──────────────────────────────────\n",
    "if PROMPT_LOG:\n",
    "    log_df = pd.DataFrame(PROMPT_LOG)\n",
    "    log_df.to_csv(\"day3_guided_lab_log.csv\", index=False)\n",
    "    print(f\"Exported {len(PROMPT_LOG)} API calls to day3_guided_lab_log.csv\")\n",
    "    print(log_df[[\"label\", \"type\", \"latency_s\"]].to_string(index=False))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Next Steps\n",
    "\n",
    "In the **Independent Lab**, you will:\n",
    "- Choose a business domain (HR, Product Docs, Support, or Research)\n",
    "- Build your own RAG knowledge base\n",
    "- Create a golden test set of 10+ questions\n",
    "- Iterate your prompt (v1 → v2) with measured improvement\n",
    "- Write an error analysis\n",
    "\n",
    "→ Proceed to the [Independent Lab](03-03_lab_2.qmd)"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}