digichat / chatbot.py
chrizefan's picture
Upload folder using huggingface_hub
fe52ef9 verified
Raw
History Blame Contribute Delete
7.19 kB
import gradio as gr
import json
import os
from typing import List, Dict, Any, Union, Generator
from generators import chat_completion_stream
from client_utils import engine_map
from file_processors import FileProcessorFactory
import logging
import time
# Configure logging first
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
# Then get a logger
logger = logging.getLogger(__name__)
class Chatbot:
def __init__(self):
"""Initialize a new chatbot with history management."""
self._stop = False # Flag to interrupt streaming
self.engine_map = engine_map
self.content_exclusions = ["<|eos|>"]
def process_file(self, files: Union[List[str], str]) -> Union[str, List[Dict[str, Any]]]:
"""Process uploaded files and return their contents as JSON."""
if isinstance(files, str):
files = [files]
file_contents = []
for file_path in files:
try:
processor = FileProcessorFactory.get_processor(file_path)
file_contents.append(processor.process(file_path))
except Exception as e:
file_contents.append(f"Error processing file {os.path.basename(file_path)}: {str(e)}\n")
return json.dumps(file_contents)
def process_message(self, message: Union[str, dict]) -> dict:
"""Convert various message formats into a standardized message dictionary."""
if isinstance(message, dict):
if message.get("text") and message.get("files"):
return {
"role": "user",
"content": f'{message["text"]}\n\nFile Upload(s):\n{self.process_file(message["files"])}'
}
elif message.get("text"):
return {
"role": "user",
"content": message["text"]
}
elif message.get("files"):
return {
"role": "user",
"content": f'File Upload(s):\n{self.process_file(message["files"])}'
}
elif message.get("content") and message.get("role"):
if isinstance(message["content"], tuple):
return {
"role": message["role"],
"content": f'File Upload(s):\n{self.process_file(message["content"][0])}'
}
elif isinstance(message["content"], str):
return {
"role": message["role"],
"content": message["content"]
}
elif isinstance(message, str):
if message:
return {
"role": "user",
"content": message,
}
return
def predict(self, message: Union[str, dict], history: List[Dict[str, Any]], engine: str, temperature: float) -> Generator:
"""Generate a streaming response for the given message and conversation history."""
self._stop = False # Reset stop flag
yield "", gr.update(interactive=False, submit_btn=False, stop_btn=True)
history.append(message)
# Process all messages and update the list
history = [msg for msg in (self.process_message(msg) for msg in history) if msg is not None]
response = []
streamed_content = "" # Accumulate the assistant's response here
def fade_html(text: str, fade_count: int = 200) -> str:
"""
Fade all characters if text is shorter than fade_count.
Otherwise, fade only the last fade_count characters.
"""
if not text:
return text
text_len = len(text)
if text_len <= fade_count:
# Fade all characters
fade_spans = ''
for i, c in enumerate(text):
opacity = 1.0 - (i / max(text_len - 1, 1))
fade_spans += f'<span style="opacity:{opacity:.2f};transition:opacity 0.2s">{c}</span>'
return fade_spans
else:
base = text[:-fade_count]
fades = text[-fade_count:]
fade_spans = ''
for i, c in enumerate(fades):
opacity = 1.0 - (i / (fade_count - 1))
fade_spans += f'<span style="opacity:{opacity:.2f};transition:opacity 0.2s">{c}</span>'
return base + fade_spans
for chunk in chat_completion_stream(history, engine=engine, temperature=temperature):
logger.debug("Chunk received: %s", chunk)
if self._stop:
break
try:
content = chunk.get("content", "")
if chunk.get("metadata"):
# Do not stream metadata contents, just update/append as before
new_metadata = True
for i, resp in enumerate(response):
if resp.get("metadata", {}).get("id") == chunk["metadata"]["id"]:
response[i]["content"] = content
response[i]["metadata"] = chunk["metadata"]
new_metadata = False
break
if new_metadata:
response.append({
"role": "assistant",
"content": content,
"metadata": chunk["metadata"]
})
# Do not stream char-by-char for metadata
# yield response, gr.skip()
else:
# Stream each character in the latest assistant message (no metadata)
if not response or response[-1].get("metadata"):
response.append({
"role": "assistant",
"content": content,
})
streamed_content = "" # Reset for new message
else:
response[-1]["content"] = content
latest = response[-1]
prev_len = len(streamed_content)
new_content = latest.get("content", "")
for c in new_content[prev_len:]:
if self._stop:
break
streamed_content += c
# Only fade the last fade_count characters of the current message
latest["content"] = fade_html(streamed_content)
yield response, gr.skip()
time.sleep(0.003)
except Exception as e:
logger.info(f"Error processing chunk from chat completion: {e}")
continue
yield response, gr.update(interactive=True, submit_btn=True, stop_btn=False)
def stop(self):
"""Interrupt the currently streaming response."""
self._stop = True