first-aid-rag-assistant / src / rag / rag_usage.py
rag_usage.py
Raw
"""RAGWithUsage — production base class that adds hybrid search and usage tracking.

Kept separate from evaluation/ so that production code does not depend on the
evaluation package. Notebooks and evaluation modules import from here too.
"""

from rag.rag_helper import RAGBase
from search import rrf, text_search, vector_search


def calc_total_price(usages):
    """Sum cost across a list of usage objects (prompt + completion tokens at $0.10/M)."""
    total = 0.0
    for usage in usages:
        total += (usage.prompt_tokens + usage.completion_tokens) * 0.1 / 1_000_000
    return total


class RAGWithUsage(RAGBase):

    def __init__(self, *args, **kwargs):
        self.text_index = kwargs.pop("text_index", None)
        self.vector_index = kwargs.pop("vector_index", None)
        self.embedder = kwargs.pop("embedder", None)

        kwargs.setdefault("index", None)
        super().__init__(*args, **kwargs)

        if self.vector_index is not None and self.embedder is not None:
            self.index = self.vector_index

        self.usages = []
        self.last_usage = None

    def reset_usage(self):
        self.usages = []
        self.last_usage = None

    def search(self, query, num_results=10):
        if (
            self.text_index is not None
            and self.vector_index is not None
            and self.embedder is not None
        ):
            text_results = text_search(query, self.text_index, num_results=10)
            vector_results = vector_search(
                query,
                self.embedder,
                self.vector_index,
                num_results=10,
            )
            return rrf(
                [text_results, vector_results],
                k=60,
                num_results=num_results,
            )

        if self.text_index is not None:
            return text_search(query, self.text_index, num_results=num_results)

        if self.vector_index is not None and self.embedder is not None:
            return vector_search(
                query,
                self.embedder,
                self.vector_index,
                num_results=num_results,
            )

    def llm(self, prompt):
        response = self._call_llm(prompt)

        self.last_usage = response.usage
        self.usages.append(response.usage)

        return response.choices[0].message.content

    def total_cost(self):
        return calc_total_price(self.usages)