RAG / app.py
Theresa
fixed indexing error, added instruction prompt to model, updated vector db creation, added input rails, made output correctness rail less strict, adapted output on failed output guardrails, added missing requirements
d69ddeb
Raw History Blame
6 kB
"""
Intelligent Knowledge Retrieval System - Minimal Frontend
"""
import streamlit as st
import time
from model.model import RAGModel
from rails import input as input_guard
from rails.output import OutputGuardrails
import secrets_local
from helper import ROLE_ASSISTANT, AUTO_ANSWERS, sanitize
from rag import retriever
from dataclasses import dataclass
from typing import List
# ============================================
# BACKEND INTEGRATION
# ============================================
@dataclass
class Answer():
answer: str
sources: List[str]
processing_time: float
def query_rag_pipeline(user_query: str, model: RAGModel, output_guardRails: OutputGuardrails, input_guardrails: input_guard.InputGuardRails, input_guardrails_active: bool = True, output_guardrails_active: bool = True) -> Answer:
"""
Query the Hugging Face model with the user query, with input and output guardrails, if enabled.
Parameters:
- user_query(str): The user input
- model (RAGModel): Model class for model interaction
- output_guardRails (OutputGuardrails): Class for output guardrails, checking for hallucinations, relevance etc.
- input_guardrails_active (bool): Whether or not to have input guard rails active
- output_guardrails_active (bool): Whether or not to have output guard rails active
Returns:
Answer
"""
start_time = time.time()
# 1. Input Guardrails
if input_guardrails_active:
checked_answer = input_guardrails.is_valid(user_query)
if not checked_answer.accepted:
return Answer(
answer=checked_answer.reason if checked_answer.reason else "Invalid input. Please try again.",
sources=[],
processing_time=start_time - time.time()
)
# 2. Context Retrieval
retrieved_docs = retriever.search(user_query)
# 3. LLM Generation
response = model.generate_response(user_query, retrieved_docs)
sources = [{"title": str(sanitize(doc))} for doc in retrieved_docs]
# 4. Output Guardrails
if output_guardrails_active:
gr_result = output_guardRails.check(user_query, response, retrieved_docs)
else:
end_time = time.time()
return Answer(
answer = sanitize(response),
sources = sources,
processing_time = end_time - start_time
)
end_time = time.time()
# 5. Final Answer
if all(gr_result.passed for gr_result in gr_result.values()):
return Answer(
answer = sanitize(response),
sources = sources,
processing_time = end_time - start_time
)
else:
return Answer(
answer = output_guardRails.format_guardrail_issues(gr_result, response),
sources = sources,
processing_time = end_time - start_time
)
# ============================================
# HAUPTANWENDUNG
# ============================================
def main():
st.set_page_config(
page_title="Knowledge Retrieval System",
page_icon="🤖",
layout="centered"
)
# Header
st.title("🤖 Intelligent Knowledge Retrieval System")
st.markdown("Stelle Fragen in natürlicher Sprache - das System durchsucht die Wissensdatenbank und generiert eine Antwort.")
st.markdown("---")
model = RAGModel(secrets_local.HF)
output_guardrails = OutputGuardrails()
input_guardrails = input_guard.InputGuardRails()
# Chat-Historie initialisieren
if "messages" not in st.session_state:
st.session_state.messages = []
# Willkommensnachricht
st.session_state.messages.append({
"role": ROLE_ASSISTANT,
"content": "Hallo! Ich kann dir bei Fragen zu unserer Wissensdatenbank helfen. Was möchtest du wissen?",
"sources":[]
})
if 'responses' not in st.session_state:
st.session_state.responses = []
# Chat-Historie anzeigen
for message in st.session_state.messages:
with st.chat_message(message["role"]):
st.write(message["content"])
# Zeige Quellen wenn vorhanden
if message["sources"]:
with st.expander("📚 Verwendete Quellen"):
for source in message["sources"]:
st.write(f"• {source['title']}")
# Chat-Eingabe
if prompt := st.chat_input("Stelle eine Frage..."):
# Benutzer-Nachricht hinzufügen und anzeigen
st.session_state.messages.append({"role": "user", "content": prompt,"sources":[]})
with st.chat_message("user"):
st.write(prompt)
# RAG-Antwort generieren
with st.chat_message(ROLE_ASSISTANT):
with st.spinner("Durchsuche Datenbank und generiere Antwort..."):
# RAG Pipeline aufrufen
print(prompt)
response = query_rag_pipeline(prompt, model, output_guardrails, input_guardrails)
# Antwort anzeigen
st.write(response.answer if response.answer else AUTO_ANSWERS.UNEXPECTED_ERROR.value)
# Antwort in Historie speichern
st.session_state.messages.append({
"role": ROLE_ASSISTANT,
"content": response.answer,
"sources": response.sources
})
# Quellen anzeigen wenn vorhanden
if response.sources:
with st.expander("📚 Verwendete Quellen"):
for source in response.sources:
st.write(f"• {source['title']}")
# Footer mit Info
st.markdown("---")
with st.expander("ℹ️ Für Entwickler"):
st.markdown("""
RAG - v1
""")
if __name__ == "__main__":
main()