| 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 |
|
|
| |
| logging.basicConfig( |
| level=logging.INFO, |
| format='%(asctime)s - %(name)s - %(levelname)s - %(message)s' |
| ) |
| |
| logger = logging.getLogger(__name__) |
|
|
| class Chatbot: |
| def __init__(self): |
| """Initialize a new chatbot with history management.""" |
| self._stop = False |
| |
| 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 |
|
|
| yield "", gr.update(interactive=False, submit_btn=False, stop_btn=True) |
|
|
| history.append(message) |
|
|
| |
| history = [msg for msg in (self.process_message(msg) for msg in history) if msg is not None] |
|
|
| response = [] |
| streamed_content = "" |
|
|
| 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_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"): |
| |
| 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"] |
| }) |
| |
| |
| else: |
| |
| if not response or response[-1].get("metadata"): |
| response.append({ |
| "role": "assistant", |
| "content": content, |
| }) |
| streamed_content = "" |
| 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 |
| |
| 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 |