File size: 2,946 Bytes
af78cff
e990dfa
 
af78cff
 
e990dfa
af78cff
e990dfa
 
 
 
 
 
 
 
af78cff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e990dfa
 
 
 
 
 
 
 
 
 
 
 
a4e6ce3
e990dfa
 
 
 
 
 
 
 
 
 
 
99d5f49
e990dfa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
99d5f49
 
 
 
 
 
 
 
 
 
 
e990dfa
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
import os
from contextlib import AsyncExitStack, asynccontextmanager

from fastapi import Request
from fastapi.responses import RedirectResponse, Response

from access_log_filter import install_access_log_filter
from editor_app import app as editor_app
from ktts_app import app as tts_app
from musicgen_app import app as music_app
from mcp_server import RouteSource, create_mcp_server
from renderer_app import app
from whisper_app import app as whisper_app


install_access_log_filter()


def _is_api_docs_path(path: str) -> bool:
    return (
        path.endswith("/openapi.json")
        or path.endswith("/docs")
        or "/docs/" in path
        or path.endswith("/redoc")
    )


@app.middleware("http")
async def protect_api_docs(request: Request, call_next):
    docs_enabled = os.getenv("ENABLE_API_DOCS", "false").lower() in {
        "1",
        "true",
        "yes",
    }
    if not docs_enabled and _is_api_docs_path(request.url.path):
        return Response(status_code=404)
    return await call_next(request)


CHILD_APPS = (
    ("/tts", tts_app, "tts"),
    ("/music", music_app, "music"),
    ("/whisper", whisper_app, "whisper"),
    ("/editor", editor_app, "editor"),
)

renderer_lifespan = app.router.lifespan_context


@app.get("/", include_in_schema=False)
async def studio_home():
    return RedirectResponse(url="/dashboard/", status_code=302)


@app.get("/apps", tags=["studio"])
async def studio_apps():
    return {
        "renderer": "/dashboard",
        "tts": "/tts/",
        "music": "/music/",
        "whisper": "/whisper/",
        "editor": "/editor/",
        "mcp": "/mcp/",
        "gradio_mcp": "/dashboard/gradio_api/mcp/",
        "mcp_tools": len(mcp_tool_catalog),
    }


for path, child_app, name in CHILD_APPS:
    app.mount(path, child_app, name=name)


mcp, mcp_tool_catalog = create_mcp_server(
    app,
    (
        RouteSource("renderer", app),
        RouteSource("tts", tts_app, "/tts"),
        RouteSource("music", music_app, "/music"),
        RouteSource("whisper", whisper_app, "/whisper"),
        RouteSource("editor", editor_app, "/editor"),
    ),
)
mcp_http_app = mcp.streamable_http_app()


@app.get("/mcp-health", include_in_schema=False)
async def mcp_health():
    return {
        "status": "ok",
        "custom_mcp_url": "/mcp/",
        "gradio_mcp_url": "/dashboard/gradio_api/mcp/",
        "custom_tool_count": len(mcp_tool_catalog),
        "transports": ["streamable-http"],
    }


@asynccontextmanager
async def studio_lifespan(root_app):
    async with AsyncExitStack() as stack:
        await stack.enter_async_context(renderer_lifespan(root_app))
        for _, child_app, _ in CHILD_APPS:
            await stack.enter_async_context(child_app.router.lifespan_context(child_app))
        await stack.enter_async_context(mcp.session_manager.run())
        yield


app.router.lifespan_context = studio_lifespan
app.mount("/mcp", mcp_http_app, name="mcp")