Spaces:
Sleeping
Sleeping
| from __future__ import annotations | |
| import os | |
| import re | |
| import stat | |
| import subprocess | |
| from datetime import datetime, timezone | |
| from pathlib import Path | |
| from threading import Lock | |
| from typing import Any | |
| from dotenv import load_dotenv | |
| from fastapi import FastAPI, File, HTTPException, Query, UploadFile | |
| from fastapi.concurrency import run_in_threadpool | |
| from fastapi.responses import Response | |
| from groq import Groq | |
| from langchain_core.messages import AIMessage, HumanMessage, SystemMessage | |
| from langchain_groq import ChatGroq | |
| from pydantic import BaseModel | |
| from fastapi.middleware.cors import CORSMiddleware | |
| load_dotenv() | |
| LANGCHAIN_PROJECT = os.getenv("LANGCHAIN_PROJECT", "Urdu Bot") | |
| os.environ["LANGCHAIN_PROJECT"] = LANGCHAIN_PROJECT | |
| os.environ["LANGSMITH_TRACING"] = os.getenv("LANGSMITH_TRACING", "true") | |
| if os.getenv("LANGSMITH_ENDPOINT"): | |
| os.environ["LANGSMITH_ENDPOINT"] = os.getenv("LANGSMITH_ENDPOINT", "") | |
| if os.getenv("LANGSMITH_API_KEY"): | |
| os.environ["LANGSMITH_API_KEY"] = os.getenv("LANGSMITH_API_KEY", "") | |
| app = FastAPI( | |
| title="Urdu Voice Chatbot Backend", | |
| description="FastAPI backend for Urdu transcription, generation, and TTS.", | |
| version="0.1.0", | |
| ) | |
| # Add this CORS configuration | |
| app.add_middleware( | |
| CORSMiddleware, | |
| allow_origins=["http://localhost:3000"], | |
| allow_credentials=True, | |
| allow_methods=["*"], | |
| allow_headers=["*"], | |
| ) | |
| session_store: dict[str, dict[str, Any]] = {} | |
| session_lock = Lock() | |
| SYSTEM_PROMPT = ( | |
| "آپ ایک مددگار اے آئی اسسٹنٹ ہیں۔ آپ صرف اردو میں جواب دیتے ہیں۔ " | |
| "اپنے جواب سادہ، رواں اور غیر رسمی مگر واضح اردو جملوں میں دیں۔ " | |
| "Markdown، bullet points، star markers، اور اردو کے علاوہ الفاظ استعمال نہ کریں۔" | |
| ) | |
| PIPER_BINARY = Path("./piper/piper") | |
| PIPER_MODEL = Path("./ur_PK-fasih-medium-model.onnx") | |
| PIPER_CONFIG = Path("./ur_PK-fasih-medium-model.onnx.json") | |
| class TranscribeResponse(BaseModel): | |
| session_id: str | |
| transcription: str | |
| turn_index: int | |
| class GenerateResponse(BaseModel): | |
| session_id: str | |
| response: str | |
| turn_index: int | |
| class SessionClearResponse(BaseModel): | |
| session_id: str | |
| cleared: bool | |
| def _utc_now() -> str: | |
| return datetime.now(timezone.utc).isoformat() | |
| def _require_groq_api_key() -> str: | |
| api_key = os.getenv("GROQ_API_KEY") | |
| if not api_key: | |
| raise HTTPException( | |
| status_code=500, | |
| detail="GROQ_API_KEY is missing. Set it in your environment before calling the API.", | |
| ) | |
| return api_key | |
| def _new_session(session_id: str) -> dict[str, Any]: | |
| timestamp = _utc_now() | |
| return { | |
| "session_id": session_id, | |
| "created_at": timestamp, | |
| "updated_at": timestamp, | |
| "inputs": [], | |
| "outputs": [], | |
| "turns": [], | |
| } | |
| def _get_or_create_session(session_id: str) -> dict[str, Any]: | |
| with session_lock: | |
| session = session_store.get(session_id) | |
| if session is None: | |
| session = _new_session(session_id) | |
| session_store[session_id] = session | |
| return session | |
| def _get_existing_session(session_id: str) -> dict[str, Any]: | |
| with session_lock: | |
| session = session_store.get(session_id) | |
| if session is None: | |
| raise HTTPException(status_code=404, detail=f"Session '{session_id}' not found.") | |
| return session | |
| def _append_transcription(session_id: str, transcription: str, source_name: str | None = None) -> int: | |
| with session_lock: | |
| session = session_store.get(session_id) | |
| if session is None: | |
| session = _new_session(session_id) | |
| session_store[session_id] = session | |
| turn = { | |
| "transcription": transcription, | |
| "response": None, | |
| "source_name": source_name, | |
| "created_at": _utc_now(), | |
| "responded_at": None, | |
| } | |
| session["inputs"].append(transcription) | |
| session["turns"].append(turn) | |
| session["updated_at"] = _utc_now() | |
| return len(session["turns"]) - 1 | |
| def _update_latest_response(session_id: str, response_text: str) -> int: | |
| with session_lock: | |
| session = session_store.get(session_id) | |
| if session is None: | |
| raise HTTPException(status_code=404, detail=f"Session '{session_id}' not found.") | |
| if not session["turns"]: | |
| raise HTTPException( | |
| status_code=409, | |
| detail="No transcription exists for this session. Call /transcribe first.", | |
| ) | |
| latest_turn = session["turns"][-1] | |
| latest_turn["response"] = response_text | |
| latest_turn["responded_at"] = _utc_now() | |
| session["outputs"].append(response_text) | |
| session["updated_at"] = _utc_now() | |
| return len(session["turns"]) - 1 | |
| def _build_history(session: dict[str, Any]) -> list[Any]: | |
| history: list[Any] = [SystemMessage(content=SYSTEM_PROMPT)] | |
| for turn in session["turns"]: | |
| transcription = turn.get("transcription") | |
| response = turn.get("response") | |
| if transcription: | |
| history.append(HumanMessage(content=transcription)) | |
| if response: | |
| history.append(AIMessage(content=response)) | |
| return history | |
| def _transcribe_audio(audio_bytes: bytes, file_name: str) -> str: | |
| api_key = _require_groq_api_key() | |
| groq_client = Groq(api_key=api_key) | |
| transcription = groq_client.audio.transcriptions.create( | |
| file=(file_name, audio_bytes), | |
| model="whisper-large-v3-turbo", | |
| language="ur", | |
| temperature=0.0, | |
| ) | |
| text = transcription.text.strip() | |
| if not text: | |
| raise HTTPException(status_code=502, detail="Transcription completed without text output.") | |
| return text | |
| def _generate_response(session_id: str) -> str: | |
| api_key = _require_groq_api_key() | |
| session = _get_existing_session(session_id) | |
| if not session["turns"]: | |
| raise HTTPException( | |
| status_code=409, | |
| detail="No transcription exists for this session. Call /transcribe first.", | |
| ) | |
| llm = ChatGroq( | |
| model="llama-3.3-70b-versatile", | |
| temperature=0.7, | |
| api_key=api_key, | |
| ) | |
| response = llm.invoke(_build_history(session)) | |
| response_text = response.content.strip() | |
| if not response_text: | |
| raise HTTPException(status_code=502, detail="LLM returned an empty response.") | |
| return response_text | |
| def normalize_tts_text(text: str) -> str: | |
| """Flatten multiline assistant text into a single pronunciation-friendly string.""" | |
| text = text.strip().strip('"“”') | |
| text = re.sub(r"(?m)^\s*[*•-]+\s*", "", text) | |
| text = re.sub(r"\n+", " ", text) | |
| text = re.sub(r"[\*•]+", "", text) | |
| text = re.sub( | |
| r"[^\u0600-\u06FF\u0750-\u077F\u08A0-\u08FF\uFB50-\uFDFF\uFE70-\uFEFF0-9A-Za-zऀ-ॿ\s،۔,:;!?()/%-]", | |
| " ", | |
| text, | |
| ) | |
| text = re.sub(r"\s+", " ", text) | |
| return text.strip() | |
| def _ensure_piper_binary() -> None: | |
| if PIPER_BINARY.exists(): | |
| current_permissions = os.stat(PIPER_BINARY).st_mode | |
| os.chmod(PIPER_BINARY, current_permissions | stat.S_IEXEC) | |
| def _synthesize_audio(text: str | None) -> bytes: | |
| if not text: | |
| raise HTTPException(status_code=409, detail="No LLM response exists for this session. Call /generate first.") | |
| if not PIPER_BINARY.exists(): | |
| raise HTTPException(status_code=500, detail="Piper binary is missing from ./piper/piper.") | |
| if not PIPER_MODEL.exists() or not PIPER_CONFIG.exists(): | |
| raise HTTPException( | |
| status_code=500, | |
| detail="Piper Urdu model files are missing from the project root.", | |
| ) | |
| _ensure_piper_binary() | |
| piper_cmd = [ | |
| str(PIPER_BINARY), | |
| "--model", | |
| str(PIPER_MODEL), | |
| "--config", | |
| str(PIPER_CONFIG), | |
| "--output_file", | |
| "-", | |
| ] | |
| try: | |
| process = subprocess.run( | |
| piper_cmd, | |
| input=normalize_tts_text(text).encode("utf-8"), | |
| stdout=subprocess.PIPE, | |
| stderr=subprocess.PIPE, | |
| check=True, | |
| ) | |
| except subprocess.CalledProcessError as exc: | |
| stderr = exc.stderr.decode("utf-8", errors="ignore") | |
| raise HTTPException(status_code=502, detail=f"TTS generation failed: {stderr}") from exc | |
| if not process.stdout: | |
| raise HTTPException(status_code=502, detail="TTS generation returned no audio output.") | |
| return process.stdout | |
| def root() -> dict[str, Any]: | |
| return { | |
| "message": "Urdu Voice Chatbot backend is running.", | |
| "endpoints": ["/transcribe", "/generate", "/tts", "/clear_session", "/health", "/docs"], | |
| } | |
| def health() -> dict[str, str]: | |
| return {"status": "ok"} | |
| async def transcribe( | |
| session_id: str = Query(..., min_length=1), | |
| audio: UploadFile = File(...), | |
| ) -> TranscribeResponse: | |
| audio_bytes = await audio.read() | |
| if not audio_bytes: | |
| raise HTTPException(status_code=400, detail="Audio file is empty.") | |
| transcription = await run_in_threadpool( | |
| _transcribe_audio, | |
| audio_bytes, | |
| audio.filename or "audio.wav", | |
| ) | |
| turn_index = _append_transcription(session_id, transcription, audio.filename) | |
| return TranscribeResponse(session_id=session_id, transcription=transcription, turn_index=turn_index) | |
| async def generate(session_id: str = Query(..., min_length=1)) -> GenerateResponse: | |
| response_text = await run_in_threadpool(_generate_response, session_id) | |
| turn_index = _update_latest_response(session_id, response_text) | |
| return GenerateResponse(session_id=session_id, response=response_text, turn_index=turn_index) | |
| async def tts(session_id: str = Query(..., min_length=1)) -> Response: | |
| session = _get_existing_session(session_id) | |
| if not session["turns"]: | |
| raise HTTPException(status_code=409, detail="No transcription exists for this session. Call /transcribe first.") | |
| latest_response = session["turns"][-1].get("response") | |
| audio_bytes = await run_in_threadpool(_synthesize_audio, latest_response) | |
| return Response( | |
| content=audio_bytes, | |
| media_type="audio/wav", | |
| headers={"Content-Disposition": f'inline; filename="{session_id}.wav"'}, | |
| ) | |
| def clear_session(session_id: str = Query(..., min_length=1)) -> SessionClearResponse: | |
| with session_lock: | |
| if session_id not in session_store: | |
| raise HTTPException(status_code=404, detail=f"Session '{session_id}' not found.") | |
| del session_store[session_id] | |
| return SessionClearResponse(session_id=session_id, cleared=True) | |