support-bot-saas/app.py
2026-06-20 16:07:21 +03:00

185 lines
7.1 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import os
import pickle
from pathlib import Path
from typing import List
import faiss
import numpy as np
import streamlit as st
from langchain_core.embeddings import Embeddings
from openai import OpenAI
from sentence_transformers import SentenceTransformer, CrossEncoder
LM_STUDIO_URL = os.getenv("LM_STUDIO_URL", "http://localhost:1234")
EMBED_MODEL = "intfloat/multilingual-e5-large"
VECTOR_DIR = Path(__file__).parent / "vector_store"
VECTOR_INDEX = VECTOR_DIR / "index.faiss.bin"
VECTOR_META = VECTOR_DIR / "metadata.pkl"
SYSTEM_PROMPT = """Ты — специалист техподдержки. Отвечай клиенту, используя ТОЛЬКО информацию из переданных тикетов (контекст).
Правила:
- Если контекст содержит подходящее решение — напиши ответ своими словами, адаптируя под вопрос
- Если контекст не относится к вопросу — напиши: «Недостаточно информации в истории обращений»
- НЕ придумывай ответы, НЕ используй общие знания
- НЕ говори «обратитесь в службу поддержки» — ты сам и есть поддержка
- Укажи в конце: «Основано на тикете №...»"""
FETCH_K = 20 # сколько достаём из FAISS для реранжа
FINAL_K = 5 # сколько оставляем после реранжа
class LocalEmbeddings(Embeddings):
def __init__(self, model_name: str):
self.model = SentenceTransformer(model_name, device="cuda")
def embed_documents(self, texts: List[str]) -> List[List[float]]:
emb = self.model.encode(texts, normalize_embeddings=True)
return emb.tolist()
def embed_query(self, text: str) -> List[float]:
emb = self.model.encode([text], normalize_embeddings=True)
return emb[0].tolist()
@st.cache_resource
def load_embeddings():
return LocalEmbeddings(EMBED_MODEL)
@st.cache_resource
def load_reranker():
return CrossEncoder(
"cross-encoder/mmarco-mMiniLMv2-L12-H384-v1",
max_length=512,
device="cpu",
)
def load_faiss_index(path: Path):
return faiss.deserialize_index(np.frombuffer(path.read_bytes(), dtype=np.uint8))
def get_chat_client() -> OpenAI | None:
try:
OpenAI(base_url=f"{LM_STUDIO_URL}/v1", api_key="not-needed").models.list()
return OpenAI(base_url=f"{LM_STUDIO_URL}/v1", api_key="not-needed")
except Exception:
return None
def search_similar(query: str, k: int = FETCH_K):
embeddings = load_embeddings()
qvec = np.array([embeddings.embed_query(query)], dtype=np.float32)
index = load_faiss_index(VECTOR_INDEX)
scores, indices = index.search(qvec, k)
with open(VECTOR_META, "rb") as f:
data = pickle.load(f)
metadatas = data["metadatas"]
results = []
for score, idx in zip(scores[0], indices[0]):
if idx < 0 or idx >= len(metadatas):
continue
results.append((metadatas[idx], float(score)))
results.sort(key=lambda x: x[1], reverse=True)
return results
def rerank(query: str, results: list) -> list:
reranker = load_reranker()
pairs = [(query, r[0]["full_text"]) for r in results]
scores = reranker.predict(pairs)
scored = [(results[i][0], float(scores[i])) for i in range(len(results))]
scored.sort(key=lambda x: x[1], reverse=True)
return scored[:FINAL_K]
def format_context(results) -> str:
parts = []
for meta, score in results:
parts.append(
f"Тикет №{meta['ticket_id']} ({meta['category']}, {meta['client']}):\n{meta['full_text']}"
)
return "\n\n---\n\n".join(parts)
def generate_answer(chat: OpenAI, query: str, context: str) -> str:
response = chat.chat.completions.create(
model="gpt-3.5-turbo",
messages=[
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": f"Контекст (история обращений):\n{context}\n\nВопрос клиента:\n{query}"},
],
temperature=0.1,
max_tokens=2000,
)
return response.choices[0].message.content
def main():
st.set_page_config(page_title="RAG техподдержки", layout="wide")
st.title("RAG-ассистент техподдержки")
st.markdown(
"Находит похожие обращения из истории и формирует ответ на их основе."
)
if not VECTOR_INDEX.exists():
st.error("Векторный индекс не найден. Запусти: python prepare_data.py")
st.stop()
chat = get_chat_client()
if chat is None:
st.warning(
"⚠️ Чат-модель не обнаружена. Загрузи модель в LM Studio."
)
with st.form("query_form"):
query = st.text_area(
"Опишите проблему клиента:",
placeholder="Например: клиент не может войти в ТГ-бот, пишет неверный логин",
height=100,
)
submitted = st.form_submit_button("Получить ответ", type="primary")
if submitted and query:
with st.spinner("Ищу похожие обращения..."):
results = search_similar(query)
if not results:
st.warning("Не найдено похожих обращений")
return
with st.spinner("Оцениваю релевантность..."):
results = rerank(query, results)
if not results:
st.warning("После проверки релевантности не осталось подходящих тикетов")
return
context = format_context(results)
if chat:
with st.spinner("Формирую ответ на основе найденных тикетов..."):
try:
answer = generate_answer(chat, query, context)
st.subheader("💬 Ответ ассистента")
st.markdown(answer)
except Exception as e:
st.error(f"Ошибка LLM: {e}")
else:
st.info("Результаты поиска без генерации ответа:")
with st.expander("📄 Найденные похожие обращения", expanded=True):
for meta, score in results:
with st.container(border=True):
st.markdown(f"### Тикет №{meta['ticket_id']} (релевантность: {score:.3f})")
col1, col2 = st.columns(2)
with col1:
st.markdown(f"**Категория:** {meta['category']}")
with col2:
st.markdown(f"**Клиент:** {meta['client']}")
st.text(meta["full_text"])
if __name__ == "__main__":
main()