| |
|
|
| import re |
| import os |
| from typing import Generator |
| from dotenv import load_dotenv |
|
|
| load_dotenv() |
|
|
| HF_TOKEN = os.getenv("HF_TOKEN", "") |
|
|
| |
| MAX_HISTORY_TURNS = 5 |
|
|
| _DEFAULT_MODEL = "meta-llama/Llama-3.1-8B-Instruct" |
|
|
| _MODEL_MAP = { |
| "qwen3:14b": _DEFAULT_MODEL, |
| "llama3.1:8b": _DEFAULT_MODEL, |
| "mistral:7b": _DEFAULT_MODEL, |
| "HuggingFaceH4/zephyr-7b-beta": _DEFAULT_MODEL, |
| "mistralai/Mistral-7B-Instruct-v0.3": _DEFAULT_MODEL, |
| "mistralai/Mistral-7B-Instruct-v0.1": _DEFAULT_MODEL, |
| "Qwen/Qwen2.5-7B-Instruct": _DEFAULT_MODEL, |
| "microsoft/Phi-3.5-mini-instruct": _DEFAULT_MODEL, |
| } |
|
|
|
|
| class ChatGenerator: |
| def __init__(self, model=None): |
| if not HF_TOKEN: |
| raise EnvironmentError( |
| "HF_TOKEN secret is missing.\n" |
| "Go to: Space β Settings β Repository secrets β add HF_TOKEN.\n" |
| "Create a token at huggingface.co/settings/tokens with " |
| "'Make calls to Inference Providers' permission." |
| ) |
|
|
| raw = model or os.getenv("HF_LLM_MODEL", _DEFAULT_MODEL) |
| self.model = _MODEL_MAP.get(raw, _DEFAULT_MODEL) |
|
|
| |
| |
| self._conversation_history: list[dict] = [] |
|
|
| from huggingface_hub import InferenceClient |
| self._client = InferenceClient( |
| provider="auto", |
| api_key=HF_TOKEN, |
| ) |
|
|
| |
|
|
| def _build_messages(self, system: str, user: str) -> list[dict]: |
| """ |
| Build the full messages list with conversation history injected: |
| [system] |
| [prior user/assistant pairs, capped to MAX_HISTORY_TURNS] |
| [current user message] |
| """ |
| messages: list[dict] = [{"role": "system", "content": system}] |
|
|
| history = self._conversation_history or [] |
| for turn in history[-MAX_HISTORY_TURNS:]: |
| if turn.get("user"): |
| messages.append({"role": "user", "content": turn["user"]}) |
| assistant_text = turn.get("assistant", "") |
| |
| if assistant_text and assistant_text != "β¦": |
| messages.append({"role": "assistant", "content": assistant_text}) |
|
|
| messages.append({"role": "user", "content": user}) |
| return messages |
|
|
| def _strip_think_tags(self, text: str) -> str: |
| return re.sub(r"<think>.*?</think>", "", text, flags=re.DOTALL).strip() |
|
|
| def _build_rag_prompts(self, query: str, relevant_docs: list) -> tuple[str, str]: |
| """Build system + user prompts from retrieved documents.""" |
| context_parts = [] |
| source_labels = [] |
|
|
| for i, (doc, score) in enumerate(relevant_docs, start=1): |
| meta = doc.metadata |
| filename = meta.get("filename") or meta.get("source") or f"Document {i}" |
| chunk_id = meta.get("chunk_id", "") |
| sfx = f" Β· chunk {chunk_id}" if chunk_id != "" else "" |
| label = f"Source {i}" |
| source_labels.append(label) |
| context_parts.append(f"[{label}: {filename}{sfx}]\n{doc.page_content}") |
|
|
| context = "\n\n---\n\n".join(context_parts) |
| source_list = ", ".join(source_labels) |
|
|
| system = f"""You are a document Q&A assistant. Answer ONLY from the CONTEXT PASSAGES below. |
| |
| RULES: |
| 1. Use ONLY information from the CONTEXT PASSAGES. No outside knowledge. |
| 2. Valid citation labels: {source_list} |
| 3. Cite inline as (Source 1) or (Source 2). No other format. |
| 4. Never invent section names, page numbers, or labels not listed above. |
| 5. If the answer is not in the context, say: "This information is not in the uploaded documents." |
| 6. You may use prior conversation turns to understand follow-up questions, but only |
| draw factual answers from the CONTEXT PASSAGES. |
| |
| CONTEXT PASSAGES: |
| {context}""" |
|
|
| user = f"Answer using ONLY the context. Cite with: {source_list}\n\nQuestion: {query}" |
| return system, user |
|
|
| |
|
|
| def _call_hf(self, system: str, user: str) -> str: |
| """Non-streaming completion. Injects conversation history automatically.""" |
| messages = self._build_messages(system, user) |
| try: |
| completion = self._client.chat.completions.create( |
| model=self.model, |
| messages=messages, |
| max_tokens=1024, |
| temperature=0.0, |
| ) |
| answer = completion.choices[0].message.content or "" |
| except Exception as exc: |
| raise RuntimeError(f"HF Inference API call failed: {exc}") from exc |
|
|
| return self._strip_think_tags(answer) |
|
|
| def generate_answer(self, query: str, relevant_docs: list) -> str: |
| """ |
| Return the full answer as a string (non-streaming). |
| |
| Conversation memory comes from self._conversation_history, which |
| app.py sets before each call: |
| generator._conversation_history = history_for_context |
| """ |
| if not relevant_docs: |
| return "No relevant documents were found. Please upload and process documents first." |
|
|
| system, user = self._build_rag_prompts(query, relevant_docs) |
| return self._call_hf(system=system, user=user) |
|
|
| |
|
|
| def generate_answer_stream(self, query: str, relevant_docs: list) -> Generator[str, None, None]: |
| """ |
| Yield answer tokens one by one as they arrive from the API. |
| |
| Use with Streamlit's st.write_stream(): |
| answer = st.write_stream(generator.generate_answer_stream(q, docs)) |
| |
| st.write_stream() renders tokens live and returns the full string when done. |
| Conversation memory is read from self._conversation_history (same as |
| generate_answer). |
| """ |
| if not relevant_docs: |
| yield "No relevant documents were found. Please upload and process documents first." |
| return |
|
|
| system, user = self._build_rag_prompts(query, relevant_docs) |
| messages = self._build_messages(system, user) |
|
|
| try: |
| stream = self._client.chat.completions.create( |
| model=self.model, |
| messages=messages, |
| max_tokens=1024, |
| temperature=0.0, |
| stream=True, |
| ) |
| except Exception as exc: |
| yield f"β οΈ Stream failed: {exc}" |
| return |
|
|
| |
| buffer = "" |
| inside_think = False |
|
|
| for chunk in stream: |
| delta = chunk.choices[0].delta.content or "" |
| if not delta: |
| continue |
|
|
| buffer += delta |
|
|
| if not inside_think: |
| if "<think>" in buffer: |
| before, _, rest = buffer.partition("<think>") |
| if before: |
| yield before |
| buffer = rest |
| inside_think = True |
| else: |
| safe_len = max(0, len(buffer) - len("<think>")) |
| if safe_len > 0: |
| yield buffer[:safe_len] |
| buffer = buffer[safe_len:] |
| else: |
| if "</think>" in buffer: |
| _, _, rest = buffer.partition("</think>") |
| buffer = rest |
| inside_think = False |
| else: |
| buffer = "" |
|
|
| |
| if buffer and not inside_think: |
| yield buffer.strip() |
|
|