wikananda's picture
Deploy Sona backend Docker Space
eebccca
Raw
History Blame Contribute Delete
2.91 kB
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from sona_ai.core import load_config, setup_logging
from sona_ai.db import init_db
from sona_ai.db.engine import SessionLocal
from sona_ai.db.models import Recording, RecordingStatus
from sona_ai.pipelines import build_speech_pipeline
from sona_ai.services import SummarizationService, TranscriptionService
from sona_ai.api.routes.projects import router as projects_router
from sona_ai.api.routes.runtime import router as runtime_router
from sona_ai.api.routes.transcribe import router as transcribe_router
from sona_ai.api.routes.summarize import router as summarize_router
from sona_ai.api.routes.chat import router as chat_router
import os
logger = setup_logging()
app = FastAPI(title="Sona AI API")
# Allow frontend to talk to this API
app.add_middleware(
CORSMiddleware,
allow_origins=["*"], # To be replaced later in production with our URL.
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
app.include_router(transcribe_router)
app.include_router(summarize_router)
app.include_router(projects_router)
app.include_router(runtime_router)
app.include_router(chat_router)
@app.on_event("startup")
async def startup_event():
logger.info("Initializing database...")
init_db()
_mark_interrupted_recordings_failed()
logger.info("Setting up environment...")
speech_config = load_config(os.getenv("SONA_SPEECH_CONFIG", "speech"))
logger.info("Loading models...")
speech_pipeline = build_speech_pipeline(speech_config)
speech_pipeline.load_models()
app.state.transcription_service = TranscriptionService(
speech_pipeline,
speech_config=speech_config,
default_model=speech_config.get("transcription", {}).get("engine", "parakeet"),
)
app.state.summarization_service = SummarizationService(
config=speech_config.get("summarization", {}).get("config", "llama"),
use_pretrained=True,
device="auto",
)
logger.info("Speech models loaded. Summarization model will load on first use.")
@app.on_event("shutdown")
async def shutdown_event():
logger.info("Shutting down...")
logger.info("Cleaning up models...")
app.state.transcription_service.close()
app.state.summarization_service.close()
logger.info("Cleanup complete!")
def _mark_interrupted_recordings_failed():
db = SessionLocal()
try:
recordings = (
db.query(Recording)
.filter(Recording.status == RecordingStatus.PROCESSING)
.all()
)
for recording in recordings:
recording.status = RecordingStatus.FAILED
recording.processing_stage = "failed"
recording.processing_job_id = None
recording.error = "Interrupted by server restart"
db.commit()
finally:
db.close()