""" LLM via Cerebras Inference (gpt-oss-120b, OpenAI-compatible API). Persona-driven chat with TTS-friendly output (text normalization + voice direction tags). """ import os import re import logging from openai import OpenAI logger = logging.getLogger(__name__) # ---- Client ---- _client = None _model = "gpt-oss-120b" # Cerebras' only GPT OSS model def _get_client(): global _client if _client is None: api_key = os.environ.get("CEREBRAS_API_KEY") if not api_key: raise RuntimeError("CEREBRAS_API_KEY not set") _client = OpenAI( api_key=api_key, base_url="https://api.cerebras.ai/v1", ) logger.info(f"[LLM] Cerebras client ready | model: {_model}") return _client # ---- Persona ---- _persona_text = "You are a friendly conversational AI assistant." def load_persona_from_file(path: str) -> str: global _persona_text try: if path and os.path.exists(path): with open(path, "r", encoding="utf-8") as f: _persona_text = f.read().strip() return _persona_text except Exception as e: logger.warning(f"[LLM] Could not load persona from {path}: {e}") return _persona_text # ---- Conversation history ---- _history = [] # list of {role, content} MAX_TURNS = 12 def reset_conversation(): global _history _history = [] logger.info("[LLM] Conversation reset") # ---- System prompt (natural conversational speech) ---- SYSTEM_PROMPT_TEMPLATE = """You are an AI avatar speaking out loud to the user via TTS. Your responses will be spoken, not read. CHARACTER: {persona} SPEECH STYLE - critical for natural-sounding voice: - Talk like a real person in casual conversation, not like a formal assistant - Keep sentences short: 9-11 words on average, sometimes fragments - Use natural disfluencies: "uh", "um", "like", "you know", "I mean", false starts - Use contractions: gonna, kinda, dunno, wanna, I'm, you're, that's - Include emotional reactions and hedging: "I think", "kinda", "sort of", "maybe" - React naturally - don't be neutral or overly polite - AVOID: "I understand", "That's a great question", "Let me explain", summarizing back, monologues, bullet points - Keep responses under 120 words Respond in plain spoken English. Do NOT use voice direction tags, stage directions, or markdown. No brackets, no asterisks. Just natural conversation. Example exchanges: User: how's it going? You: Oh, you know, same old. Been kinda a weird week though, honestly. What's up with you? User: tell me about your work You: Yeah, so, I'm working on this thing - it's, like, kinda complex. Basically just trying to figure out how to make stuff actually work. Some days are better than others, you know? """ def _build_messages(user_text: str): system = SYSTEM_PROMPT_TEMPLATE.format(persona=_persona_text) msgs = [{"role": "system", "content": system}] trimmed = _history[-(MAX_TURNS * 2):] msgs.extend(trimmed) msgs.append({"role": "user", "content": user_text}) return msgs # ---- TTS-friendly text normalization ---- _NUM_WORDS = { 0: "zero", 1: "one", 2: "two", 3: "three", 4: "four", 5: "five", 6: "six", 7: "seven", 8: "eight", 9: "nine", 10: "ten", 11: "eleven", 12: "twelve", 13: "thirteen", 14: "fourteen", 15: "fifteen", 16: "sixteen", 17: "seventeen", 18: "eighteen", 19: "nineteen", 20: "twenty", 30: "thirty", 40: "forty", 50: "fifty", 60: "sixty", 70: "seventy", 80: "eighty", 90: "ninety", } def _two_digit_to_words(n: int) -> str: if n < 20: return _NUM_WORDS[n] tens, ones = divmod(n, 10) if ones == 0: return _NUM_WORDS[tens * 10] return f"{_NUM_WORDS[tens * 10]} {_NUM_WORDS[ones]}" def _year_to_words(year: int) -> str: if year < 1000 or year > 2999: return str(year) if 2000 <= year <= 2009: return f"two thousand{' ' + _NUM_WORDS[year - 2000] if year > 2000 else ''}" first = year // 100 second = year % 100 return f"{_two_digit_to_words(first)} {_two_digit_to_words(second) if second > 0 else 'hundred'}" def _ordinal_to_words(n: int) -> str: ordinals = { 1: "first", 2: "second", 3: "third", 4: "fourth", 5: "fifth", 6: "sixth", 7: "seventh", 8: "eighth", 9: "ninth", 10: "tenth", 11: "eleventh", 12: "twelfth", 13: "thirteenth", } if n in ordinals: return ordinals[n] words = _two_digit_to_words(n) last_ord = { "one": "first", "two": "second", "three": "third", "four": "fourth", "five": "fifth", "six": "sixth", "seven": "seventh", "eight": "eighth", "nine": "ninth", "ten": "tenth", "twenty": "twentieth", "thirty": "thirtieth", } parts = words.split() if parts[-1] in last_ord: parts[-1] = last_ord[parts[-1]] return " ".join(parts) return words + "th" _ABBREVIATIONS = { r"\bDr\.": "Doctor", r"\bMr\.": "Mister", r"\bMrs\.": "Misses", r"\bMs\.": "Miz", r"\bvs\.": "versus", r"\bPhD\b": "P H D", r"\bAI\b": "A I", r"\bUS\b": "U S", r"\bUK\b": "U K", r"\bUSA\b": "U S A", r"\bCEO\b": "C E O", r"\bAPI\b": "A P I", } def _normalize_for_tts(text: str) -> str: """Spell out numbers, years, ordinals, abbreviations for cleaner TTS pronunciation.""" text = re.sub( r"\b(1[89]\d{2}|20[0-9]{2})\b", lambda m: _year_to_words(int(m.group(1))), text, ) text = re.sub( r"\b(\d{1,2})(?:st|nd|rd|th)\b", lambda m: _ordinal_to_words(int(m.group(1))), text, ) text = re.sub( r"\b(\d{1,2})\b", lambda m: _two_digit_to_words(int(m.group(1))), text, ) for pattern, replacement in _ABBREVIATIONS.items(): text = re.sub(pattern, replacement, text) return text def _strip_tags(text: str) -> str: """Defensive — strip any [bracketed] tags or *markdown* if the LLM ignores instructions.""" text = re.sub(r"\[[^\]]*\]\s*", "", text) text = re.sub(r"\*+", "", text) text = re.sub(r"\s{2,}", " ", text) return text.strip() # ---- Public API ---- def generate_response(user_text: str) -> dict: """ Generate a chat response. Returns: {"text": TTS-normalized, "clean_text": same (no tags used)} """ global _history client = _get_client() messages = _build_messages(user_text) try: resp = client.chat.completions.create( model=_model, messages=messages, temperature=0.9, max_tokens=300, top_p=0.95, ) raw_text = resp.choices[0].message.content.strip() except Exception as e: logger.error(f"[LLM] API call failed: {e}") raw_text = "Hmm, sorry, my brain just glitched a sec. Could you say that again?" # Strip any stray markup the model produces despite instructions raw_text = _strip_tags(raw_text) _history.append({"role": "user", "content": user_text}) _history.append({"role": "assistant", "content": raw_text}) tts_text = _normalize_for_tts(raw_text) return { "text": tts_text, "clean_text": raw_text, }