aboutsummaryrefslogtreecommitdiff
path: root/app/backends/probe.py
diff options
context:
space:
mode:
Diffstat (limited to 'app/backends/probe.py')
-rw-r--r--app/backends/probe.py128
1 files changed, 128 insertions, 0 deletions
diff --git a/app/backends/probe.py b/app/backends/probe.py
new file mode 100644
index 0000000..ebef86e
--- /dev/null
+++ b/app/backends/probe.py
@@ -0,0 +1,128 @@
+"""Identify which TTS backend answers at a URL (remote-server probing).
+
+The hub keeps locally-managed backends distinct from externally-run ones: a
+server this tool started is tagged "[local]", and a server found by probing a
+configured remote URL (``*_REMOTE_URL`` in ``app/converter/config.py``) is
+tagged "[remote]". To know that a remote URL really hosts the backend we
+think it does (and not some other HTTP service), each backend exposes a small
+identity check over plain HTTP:
+
+ * audio.cpp ``GET /health`` -> ``{"status": "ok"}`` and ``GET /v1/models``
+ -> ``{"data": [{"id": ...}, ...]}``.
+ * faster ``GET /health`` -> a JSON object with a ``model_loaded`` key.
+ * qwen-tts a Gradio app: ``GET /info`` -> ``named_endpoints`` containing
+ the endpoint names the converter calls (``/run_instruct`` /
+ ``/run_custom_voice`` / ``/generate_custom_voice`` for the
+ CustomVoice demo; ``/run_voice_clone`` / ``/generate_voice_clone``
+ for the Base demo).
+
+``identify_server`` returns one of the IDENTITY_* constants, or None when the
+URL does not answer or answers as something unrecognized. It is stdlib-only
+(urllib) and deliberately imports nothing from the other backend modules, so
+it stays cheap to import alongside ``backends.common``.
+"""
+
+import json
+import urllib.parse
+import urllib.request
+from typing import Optional
+
+from backends import common
+
+IDENTITY_AUDIOCPP = "audiocpp"
+IDENTITY_FASTER = "faster"
+IDENTITY_QWEN_CUSTOM = "qwen-custom"
+IDENTITY_QWEN_CLONE = "qwen-clone"
+
+# Endpoint names the converter resolves for each qwen demo server (see
+# converter.tts QwenTTSClient). Mirror them here so identification matches
+# exactly what the converter would call.
+_QWEN_CUSTOM_ENDPOINTS = (
+ "/run_instruct", "/run_custom_voice", "/generate_custom_voice")
+_QWEN_CLONE_ENDPOINTS = ("/run_voice_clone", "/generate_voice_clone")
+
+DEFAULT_TIMEOUT = 3.0
+
+
+def identify_server(url: str, timeout: float = DEFAULT_TIMEOUT) -> Optional[str]:
+ """Return the backend identity answering at URL, or None.
+
+ A cheap TCP-connect gate runs first (``common.server_running``) so a dead
+ or unrouteable host returns quickly; the HTTP probes only run when
+ something is listening. Returns None when the URL is empty/unparsable,
+ unreachable, or answers as none of the known backends.
+ """
+ if not url:
+ return None
+ base = url.rstrip("/")
+ if not common.server_running(url):
+ return None
+ identity = _identify_health(base, timeout)
+ if identity is not None:
+ return identity
+ return _identify_gradio(base, timeout)
+
+
+def _get_json(url: str, timeout: float) -> Optional[dict]:
+ """GET URL and parse a JSON object, or None on any error."""
+ try:
+ with urllib.request.urlopen(url, timeout=timeout) as response:
+ payload = json.loads(response.read().decode("utf-8"))
+ except (OSError, ValueError):
+ return None
+ return payload if isinstance(payload, dict) else None
+
+
+def _identify_health(base: str, timeout: float) -> Optional[str]:
+ """Identify audio.cpp / faster from their ``/health`` responses."""
+ payload = _get_json(f"{base}/health", timeout)
+ if payload is None:
+ return None
+ # faster's /health reports model load state under "model_loaded".
+ if "model_loaded" in payload:
+ return IDENTITY_FASTER
+ # audio.cpp's /health reports {"status": "ok"}; confirm it also serves
+ # the /v1/models catalog (id-bearing entries) to avoid mistaking some
+ # other service that happens to return {"status": "ok"}.
+ if payload.get("status") == "ok":
+ models = _get_json(f"{base}/v1/models", timeout)
+ entries = models.get("data") if models is not None else None
+ if isinstance(entries, list) and entries \
+ and any(isinstance(e, dict) and e.get("id") for e in entries):
+ return IDENTITY_AUDIOCPP
+ return None
+
+
+def _identify_gradio(base: str, timeout: float) -> Optional[str]:
+ """Identify a qwen-tts Gradio demo from its ``/info`` named endpoints."""
+ payload = _get_json(f"{base}/info", timeout)
+ if payload is None:
+ return None
+ endpoints = payload.get("named_endpoints")
+ if not isinstance(endpoints, dict):
+ return None
+ if any(name in endpoints for name in _QWEN_CUSTOM_ENDPOINTS):
+ return IDENTITY_QWEN_CUSTOM
+ if any(name in endpoints for name in _QWEN_CLONE_ENDPOINTS):
+ return IDENTITY_QWEN_CLONE
+ return None
+
+
+def same_endpoint(url_a: str, url_b: str) -> bool:
+ """True when URL_A and URL_B address the same host and port.
+
+ Scheme and path are ignored (127.0.0.1:8080 and http://127.0.0.1:8080/
+ are the same server). Returns False when either URL is empty/unparsable.
+ """
+ if not url_a or not url_b:
+ return False
+ try:
+ a = urllib.parse.urlsplit(url_a)
+ b = urllib.parse.urlsplit(url_b)
+ except ValueError:
+ return False
+ host_a = a.hostname or "127.0.0.1"
+ host_b = b.hostname or "127.0.0.1"
+ port_a = a.port or (443 if (a.scheme or "http") == "https" else 80)
+ port_b = b.port or (443 if (b.scheme or "http") == "https" else 80)
+ return host_a == host_b and port_a == port_b