"""rag_demo.py — retriever + prompt + ChatOpenAI"""

import shutil
from pathlib import Path

from dotenv import load_dotenv
from langchain_community.vectorstores import Chroma
from langchain_core.output_parsers import StrOutputParser
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.runnables import RunnablePassthrough
from langchain_openai import ChatOpenAI, OpenAIEmbeddings

load_dotenv()

TEXTS = [
    "The <a> tag creates a hyperlink. Use the href attribute to set the destination URL.",
    "The <title> tag sets the page title shown in the browser tab.",
    "The <h1> tag marks the main heading on a page.",
]

QUESTION = "How do I make a link on a page?"
DB_DIR = Path(__file__).parent / "rag_chroma_db"


def format_docs(docs: list) -> str:
    return "\n\n".join(doc.page_content for doc in docs)


def main() -> None:
    if DB_DIR.exists():
        shutil.rmtree(DB_DIR)

    embeddings = OpenAIEmbeddings(model="text-embedding-3-small")
    vectorstore = Chroma.from_texts(
        texts=TEXTS,
        embedding=embeddings,
        persist_directory=str(DB_DIR),
    )

    retriever = vectorstore.as_retriever(search_kwargs={"k": 2})

    print("=== Retrieved chunks ===")
    docs = retriever.invoke(QUESTION)
    for i, doc in enumerate(docs):
        print(f"[{i}] {doc.page_content}")

    prompt = ChatPromptTemplate.from_template(
        "Answer in one short sentence using only the context below.\n\n"
        "Context:\n{context}\n\n"
        "Question: {question}"
    )

    llm = ChatOpenAI(model="gpt-4o-mini", temperature=0)

    rag_chain = (
        {"context": retriever | format_docs, "question": RunnablePassthrough()}
        | prompt
        | llm
        | StrOutputParser()
    )

    print("\n=== Answer ===")
    print(f"Q: {QUESTION}")
    answer = rag_chain.invoke(QUESTION)
    print(f"A: {answer}")


if __name__ == "__main__":
    main()