vrex-flow-engine 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,125 @@
1
+ """Live status dashboard (rich) for the flow-engine supervisor.
2
+
3
+ Read-only view rendered from `UiState` every refresh: engine + tunnel health,
4
+ the irreplaceable 'extension connected' signal, and a merged log tail. Ctrl-C in
5
+ the supervisor tears it down cleanly.
6
+ """
7
+ from __future__ import annotations
8
+
9
+ from collections import deque
10
+ from dataclasses import dataclass, field
11
+ from typing import Deque, Optional, Tuple
12
+
13
+ from rich.columns import Columns
14
+ from rich.console import Group
15
+ from rich.panel import Panel
16
+ from rich.table import Table
17
+ from rich.text import Text
18
+
19
+ from .config import DEFAULT_PUBLIC_BASE_URL
20
+
21
+ # (source, text) log entries — source is the process name ("engine"/"tunnel").
22
+ LogEntry = Tuple[str, str]
23
+
24
+
25
+ @dataclass
26
+ class UiState:
27
+ port: int = 8101
28
+ public_base_url: str = DEFAULT_PUBLIC_BASE_URL
29
+ with_tunnel: bool = True
30
+ uptime_s: float = 0.0
31
+ # engine
32
+ engine_alive: bool = False
33
+ engine_pid: Optional[int] = None
34
+ engine_restarts: int = 0
35
+ health: Optional[dict] = None
36
+ # tunnel
37
+ tunnel_alive: bool = False
38
+ tunnel_pid: Optional[int] = None
39
+ tunnel_restarts: int = 0
40
+ tunnel_connected: bool = False
41
+ logs: Deque[LogEntry] = field(default_factory=deque)
42
+
43
+
44
+ def _fmt_uptime(secs: float) -> str:
45
+ s = int(secs)
46
+ h, s = divmod(s, 3600)
47
+ m, s = divmod(s, 60)
48
+ return f"{h}h{m:02d}m{s:02d}s" if h else f"{m}m{s:02d}s"
49
+
50
+
51
+ def _dot(ok: bool, warn: bool = False) -> Text:
52
+ if warn:
53
+ return Text("●", style="yellow")
54
+ return Text("●", style="green" if ok else "red")
55
+
56
+
57
+ def _engine_panel(st: UiState) -> Panel:
58
+ h = st.health or {}
59
+ ext_connected = bool(h.get("extension_connected"))
60
+ t = Table.grid(padding=(0, 1))
61
+ t.add_column(justify="right", style="dim")
62
+ t.add_column()
63
+ status = Text.assemble(_dot(st.engine_alive), " ", ("UP" if st.engine_alive else "DOWN"))
64
+ t.add_row("status", status)
65
+ t.add_row("health", Text("responding", style="green") if st.health else Text("no response", style="red"))
66
+ t.add_row("port", str(st.port))
67
+ t.add_row("pid", str(st.engine_pid or "—"))
68
+ t.add_row(
69
+ "extension",
70
+ Text.assemble(_dot(ext_connected, warn=not ext_connected), " ",
71
+ "connected" if ext_connected else "OPEN A FLOW TAB"),
72
+ )
73
+ t.add_row("instances", str(h.get("instances", "—")))
74
+ if st.engine_restarts:
75
+ t.add_row("restarts", Text(str(st.engine_restarts), style="yellow"))
76
+ return Panel(t, title="[bold]engine[/]", border_style="green" if st.engine_alive else "red")
77
+
78
+
79
+ def _tunnel_panel(st: UiState) -> Panel:
80
+ if not st.with_tunnel:
81
+ return Panel(Text("disabled (--no-tunnel)", style="dim"), title="[bold]tunnel[/]", border_style="dim")
82
+ connected = st.tunnel_connected and st.tunnel_alive
83
+ t = Table.grid(padding=(0, 1))
84
+ t.add_column(justify="right", style="dim")
85
+ t.add_column()
86
+ status = Text.assemble(_dot(connected, warn=st.tunnel_alive and not connected), " ",
87
+ "connected" if connected else ("connecting…" if st.tunnel_alive else "DOWN"))
88
+ t.add_row("status", status)
89
+ host = st.public_base_url.replace("https://", "").replace("http://", "")
90
+ t.add_row("public", Text(host, style="cyan"))
91
+ t.add_row("→", f"127.0.0.1:{st.port}")
92
+ t.add_row("pid", str(st.tunnel_pid or "—"))
93
+ if st.tunnel_restarts:
94
+ t.add_row("restarts", Text(str(st.tunnel_restarts), style="yellow"))
95
+ return Panel(t, title="[bold]tunnel[/]", border_style="green" if connected else "yellow")
96
+
97
+
98
+ def _logs_panel(st: UiState, limit: int = 14) -> Panel:
99
+ body = Text()
100
+ entries = list(st.logs)[-limit:] if st.logs else []
101
+ if not entries:
102
+ body.append("waiting for output…", style="dim")
103
+ for source, line in entries:
104
+ tag_style = "cyan" if source == "engine" else "magenta"
105
+ body.append(f"{source:>6} ", style=tag_style)
106
+ body.append(line[:200] + "\n", style="default")
107
+ return Panel(body, title="[bold]logs[/] [dim](engine + tunnel)[/]", border_style="blue")
108
+
109
+
110
+ def render(st: UiState):
111
+ """Build the full dashboard renderable for one refresh."""
112
+ overall_up = st.engine_alive and (not st.with_tunnel or (st.tunnel_alive and st.tunnel_connected))
113
+ header = Panel(
114
+ Text.assemble(
115
+ ("🎬 Vrex Flow Engine", "bold"),
116
+ " ",
117
+ _dot(overall_up, warn=st.engine_alive and not overall_up),
118
+ (" running" if overall_up else " starting", "dim"),
119
+ (" uptime ", "dim"), (_fmt_uptime(st.uptime_s), "bold"),
120
+ ),
121
+ border_style="bold blue",
122
+ )
123
+ status_row = Columns([_engine_panel(st), _tunnel_panel(st)], equal=True, expand=True)
124
+ footer = Text("Ctrl-C to stop · logs at ~/.vrex-flow/logs", style="dim", justify="center")
125
+ return Group(header, status_row, _logs_panel(st), footer)
@@ -0,0 +1,40 @@
1
+ """Engine + tunnel health signals for the dashboard.
2
+
3
+ - Engine: GET /api/health on loopback → {ok, extension_connected, instances}.
4
+ - Tunnel: cloudflared's own logs are the source of truth; `tunnel_state_from_line`
5
+ classifies a log line into connected / error / None so the supervisor can track
6
+ whether the tunnel has registered a connection to Cloudflare's edge.
7
+ """
8
+ from __future__ import annotations
9
+
10
+ from typing import Optional
11
+
12
+ import httpx
13
+
14
+
15
+ def probe_engine(port: int, timeout: float = 2.0) -> Optional[dict]:
16
+ """Return the /api/health JSON, or None if the engine isn't answering yet."""
17
+ try:
18
+ resp = httpx.get(f"http://127.0.0.1:{port}/api/health", timeout=timeout)
19
+ if resp.status_code == 200:
20
+ return resp.json()
21
+ except Exception:
22
+ return None
23
+ return None
24
+
25
+
26
+ def tunnel_state_from_line(line: str) -> Optional[str]:
27
+ """Map a cloudflared log line to 'connected' | 'error' | None (no change).
28
+
29
+ cloudflared logs 'Registered tunnel connection' once an edge connection is up,
30
+ and 'Unregistered tunnel connection' / 'ERR' when it drops. We treat a
31
+ registration as connected and an explicit error/unregister as a drop.
32
+
33
+ Order matters: 'unregistered tunnel connection' CONTAINS 'registered tunnel
34
+ connection', so the drop/error cases must be tested before the connect case."""
35
+ low = line.lower()
36
+ if "unregistered tunnel connection" in low or "failed to" in low or "err " in low or "error=" in low:
37
+ return "error"
38
+ if "registered tunnel connection" in low:
39
+ return "connected"
40
+ return None
@@ -0,0 +1,109 @@
1
+ """ManagedProcess — a supervised child process with captured, tee'd logs.
2
+
3
+ The supervisor runs two of these: the uvicorn engine and the cloudflared tunnel.
4
+ Each spawns with a background thread pumping merged stdout/stderr into (a) its own
5
+ rotating-append log file and (b) an `on_line(name, text)` sink the dashboard tails.
6
+ Restart/backoff policy lives in the supervisor loop, not here — this class only
7
+ owns spawn, log capture, liveness, and graceful stop.
8
+ """
9
+ from __future__ import annotations
10
+
11
+ import os
12
+ import signal
13
+ import subprocess
14
+ import threading
15
+ import time
16
+ from pathlib import Path
17
+ from typing import Callable, Optional, Sequence
18
+
19
+ OnLine = Callable[[str, str], None]
20
+
21
+
22
+ class ManagedProcess:
23
+ def __init__(
24
+ self,
25
+ name: str,
26
+ argv: Sequence[str],
27
+ *,
28
+ env: Optional[dict] = None,
29
+ log_file: Optional[Path] = None,
30
+ on_line: Optional[OnLine] = None,
31
+ display: Optional[str] = None,
32
+ ) -> None:
33
+ self.name = name
34
+ self.argv = list(argv)
35
+ # `display` is a secret-masked argv for the UI (e.g. cloudflared --token ***).
36
+ self.display = display or " ".join(self.argv)
37
+ self._env = {**os.environ, **(env or {})}
38
+ self._log_file = log_file
39
+ self._on_line = on_line
40
+
41
+ self._proc: Optional[subprocess.Popen] = None
42
+ self._pump: Optional[threading.Thread] = None
43
+ self.started_at: Optional[float] = None
44
+ self.restarts = 0
45
+ self.last_exit: Optional[int] = None
46
+
47
+ # ── lifecycle ────────────────────────────────────────────────────────────
48
+ def start(self) -> None:
49
+ """Spawn the child (merged stderr→stdout, line-buffered) + pump thread."""
50
+ if self._log_file:
51
+ self._log_file.parent.mkdir(parents=True, exist_ok=True)
52
+ self._proc = subprocess.Popen(
53
+ self.argv,
54
+ env=self._env,
55
+ stdout=subprocess.PIPE,
56
+ stderr=subprocess.STDOUT,
57
+ text=True,
58
+ bufsize=1,
59
+ )
60
+ self.started_at = time.monotonic()
61
+ self._pump = threading.Thread(target=self._read_output, name=f"{self.name}-pump", daemon=True)
62
+ self._pump.start()
63
+
64
+ def _read_output(self) -> None:
65
+ assert self._proc and self._proc.stdout
66
+ log_fh = open(self._log_file, "a", buffering=1) if self._log_file else None
67
+ try:
68
+ for line in self._proc.stdout:
69
+ text = line.rstrip("\n")
70
+ if log_fh:
71
+ log_fh.write(text + "\n")
72
+ if self._on_line:
73
+ self._on_line(self.name, text)
74
+ finally:
75
+ if log_fh:
76
+ log_fh.close()
77
+
78
+ @property
79
+ def alive(self) -> bool:
80
+ return self._proc is not None and self._proc.poll() is None
81
+
82
+ def poll(self) -> Optional[int]:
83
+ """Return the exit code if the process has died (and record it), else None."""
84
+ if self._proc is None:
85
+ return None
86
+ code = self._proc.poll()
87
+ if code is not None and self.last_exit != code:
88
+ self.last_exit = code
89
+ return code
90
+
91
+ @property
92
+ def pid(self) -> Optional[int]:
93
+ return self._proc.pid if self._proc else None
94
+
95
+ @property
96
+ def uptime_s(self) -> float:
97
+ return (time.monotonic() - self.started_at) if self.started_at and self.alive else 0.0
98
+
99
+ def stop(self, grace: float = 8.0) -> None:
100
+ """SIGTERM, wait up to `grace`, then SIGKILL. Idempotent."""
101
+ if not self._proc or self._proc.poll() is not None:
102
+ return
103
+ try:
104
+ self._proc.send_signal(signal.SIGTERM)
105
+ self._proc.wait(timeout=grace)
106
+ except subprocess.TimeoutExpired:
107
+ self._proc.kill()
108
+ except ProcessLookupError:
109
+ pass
@@ -0,0 +1,184 @@
1
+ """Supervisor — spawn + babysit the engine and the cloudflared tunnel, and drive
2
+ the live dashboard.
3
+
4
+ Process tree: this launcher (parent) → uvicorn engine (`python -m flow_engine.main`)
5
+ + cloudflared tunnel, each a child ManagedProcess. The engine's flow/wavespeed
6
+ config is its own concern; the supervisor only injects the API key + ports and
7
+ watches liveness. Dead children restart with capped exponential backoff; Ctrl-C
8
+ or SIGTERM stops both gracefully.
9
+ """
10
+ from __future__ import annotations
11
+
12
+ import shutil
13
+ import signal
14
+ import socket
15
+ import sys
16
+ import threading
17
+ import time
18
+ from collections import deque
19
+ from typing import Optional
20
+
21
+ from rich.console import Console
22
+
23
+ from . import dashboard, health
24
+ from .config import LOGS_DIR, FlowConfig, mask
25
+ from .processes import ManagedProcess
26
+
27
+ console = Console()
28
+
29
+ STABLE_UPTIME_S = 60.0 # a child alive this long → reset its backoff counter
30
+ BASE_BACKOFF_S = 1.0
31
+ MAX_BACKOFF_S = 30.0
32
+ HEALTH_EVERY_S = 2.0
33
+
34
+
35
+ def preflight(cfg: FlowConfig, with_tunnel: bool) -> bool:
36
+ """Verify prerequisites the launcher can check. Returns False to abort."""
37
+ ok = True
38
+ if not cfg.has_engine_key:
39
+ console.print("[red]✗[/] no engine API key — run [bold]flow-engine setup[/]")
40
+ ok = False
41
+ if with_tunnel and not shutil.which("cloudflared"):
42
+ console.print(
43
+ "[red]✗[/] cloudflared not found. Install it:\n"
44
+ " macOS: [bold]brew install cloudflared[/]\n"
45
+ " other: https://developers.cloudflare.com/cloudflare-one/connections/connect-networks/downloads/"
46
+ )
47
+ ok = False
48
+ if with_tunnel and not cfg.has_tunnel_token:
49
+ console.print("[red]✗[/] no tunnel token — run [bold]flow-engine setup[/] (or pass --no-tunnel)")
50
+ ok = False
51
+ if _port_in_use(cfg.http_port):
52
+ console.print(f"[yellow]![/] port {cfg.http_port} already in use — the engine may fail to bind")
53
+ return ok
54
+
55
+
56
+ def _port_in_use(port: int) -> bool:
57
+ with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
58
+ return s.connect_ex(("127.0.0.1", port)) == 0
59
+
60
+
61
+ class Supervisor:
62
+ def __init__(self, cfg: FlowConfig, with_tunnel: bool = True) -> None:
63
+ self.cfg = cfg
64
+ self.with_tunnel = with_tunnel
65
+ self._stop = threading.Event()
66
+ self._logs: deque = deque(maxlen=500)
67
+ self._tunnel_connected = False
68
+ self._restart_at: dict[str, Optional[float]] = {}
69
+ self._started_at = time.monotonic()
70
+
71
+ def _on_line(self, source: str, text: str) -> None:
72
+ """Pump-thread callback: buffer the line + update tunnel connection state."""
73
+ self._logs.append((source, text))
74
+ if source == "tunnel":
75
+ state = health.tunnel_state_from_line(text)
76
+ if state == "connected":
77
+ self._tunnel_connected = True
78
+ elif state == "error":
79
+ self._tunnel_connected = False
80
+
81
+ def _build_engine(self) -> ManagedProcess:
82
+ env = {
83
+ "FLOW_ENGINE_API_KEY": self.cfg.engine_api_key,
84
+ "FLOWPROXY_HTTP_PORT": str(self.cfg.http_port),
85
+ "FLOWPROXY_PUBLIC_BASE_URL": self.cfg.public_base_url,
86
+ }
87
+ return ManagedProcess(
88
+ "engine",
89
+ [sys.executable, "-m", "flow_engine.main"],
90
+ env=env,
91
+ log_file=LOGS_DIR / "engine.log",
92
+ on_line=self._on_line,
93
+ display=f"{sys.executable} -m flow_engine.main",
94
+ )
95
+
96
+ def _build_tunnel(self) -> ManagedProcess:
97
+ token = self.cfg.tunnel_token
98
+ argv = ["cloudflared", "--no-autoupdate", "tunnel", "run", "--token", token]
99
+ return ManagedProcess(
100
+ "tunnel",
101
+ argv,
102
+ log_file=LOGS_DIR / "tunnel.log",
103
+ on_line=self._on_line,
104
+ display=f"cloudflared --no-autoupdate tunnel run --token {mask(token)}",
105
+ )
106
+
107
+ def _supervise(self, proc: ManagedProcess) -> None:
108
+ """Restart a dead child with capped exponential backoff."""
109
+ proc.poll() # record exit code if it just died
110
+ if proc.alive:
111
+ if proc.uptime_s > STABLE_UPTIME_S and proc.restarts:
112
+ proc.restarts = 0
113
+ self._restart_at[proc.name] = None
114
+ return
115
+ due = self._restart_at.get(proc.name)
116
+ if due is None:
117
+ proc.restarts += 1
118
+ delay = min(MAX_BACKOFF_S, BASE_BACKOFF_S * (2 ** (proc.restarts - 1)))
119
+ self._restart_at[proc.name] = time.monotonic() + delay
120
+ self._logs.append(("system", f"{proc.name} exited (code {proc.last_exit}); restart in {int(delay)}s"))
121
+ elif time.monotonic() >= due:
122
+ self._restart_at[proc.name] = None
123
+ proc.start()
124
+ self._logs.append(("system", f"{proc.name} restarted (attempt {proc.restarts})"))
125
+
126
+ def run(self) -> int:
127
+ LOGS_DIR.mkdir(parents=True, exist_ok=True)
128
+ engine = self._build_engine()
129
+ tunnel = self._build_tunnel() if self.with_tunnel else None
130
+
131
+ signal.signal(signal.SIGTERM, lambda *_: self._stop.set())
132
+
133
+ console.print(f"[dim]starting engine on :{self.cfg.http_port}"
134
+ + (f" · tunnel → {self.cfg.public_base_url}" if self.with_tunnel else " · tunnel disabled")
135
+ + "[/]")
136
+ engine.start()
137
+ if tunnel:
138
+ tunnel.start()
139
+
140
+ last_probe = 0.0
141
+ cached_health: Optional[dict] = None
142
+ from rich.live import Live
143
+
144
+ try:
145
+ with Live(console=console, refresh_per_second=8, screen=True) as live:
146
+ while not self._stop.is_set():
147
+ now = time.monotonic()
148
+ if now - last_probe >= HEALTH_EVERY_S:
149
+ cached_health = health.probe_engine(self.cfg.http_port)
150
+ last_probe = now
151
+ self._supervise(engine)
152
+ if tunnel:
153
+ self._supervise(tunnel)
154
+ live.update(dashboard.render(self._state(engine, tunnel, cached_health)))
155
+ time.sleep(0.25)
156
+ except KeyboardInterrupt:
157
+ pass
158
+ finally:
159
+ self._shutdown(engine, tunnel)
160
+ return 0
161
+
162
+ def _state(self, engine, tunnel, cached_health) -> "dashboard.UiState":
163
+ return dashboard.UiState(
164
+ port=self.cfg.http_port,
165
+ public_base_url=self.cfg.public_base_url,
166
+ with_tunnel=self.with_tunnel,
167
+ uptime_s=time.monotonic() - self._started_at,
168
+ engine_alive=engine.alive,
169
+ engine_pid=engine.pid,
170
+ engine_restarts=engine.restarts,
171
+ health=cached_health,
172
+ tunnel_alive=bool(tunnel and tunnel.alive),
173
+ tunnel_pid=tunnel.pid if tunnel else None,
174
+ tunnel_restarts=tunnel.restarts if tunnel else 0,
175
+ tunnel_connected=self._tunnel_connected,
176
+ logs=self._logs,
177
+ )
178
+
179
+ def _shutdown(self, engine, tunnel) -> None:
180
+ console.print("\n[dim]stopping…[/]")
181
+ if tunnel:
182
+ tunnel.stop()
183
+ engine.stop()
184
+ console.print("[green]✓[/] engine + tunnel stopped")
flow_engine/config.py ADDED
@@ -0,0 +1,43 @@
1
+ """Runtime configuration for flow_engine.
2
+
3
+ Ports default to the SAME values flowboard's agent uses (HTTP 8101, WS 9223)
4
+ so the existing extension build connects with zero changes. Because of that,
5
+ flow_engine and the flowboard agent cannot both run on the defaults at once —
6
+ override FLOWPROXY_HTTP_PORT / FLOWPROXY_EXT_WS_PORT (and rebuild the
7
+ extension's hardcoded URLs in extension/background.js) if you need both.
8
+ """
9
+ from __future__ import annotations
10
+
11
+ import os
12
+ from pathlib import Path
13
+
14
+ ROOT = Path(__file__).resolve().parent.parent
15
+ STORAGE_DIR = Path(os.getenv("FLOWPROXY_STORAGE", ROOT / "storage"))
16
+
17
+ HTTP_PORT = int(os.getenv("FLOWPROXY_HTTP_PORT", "8101"))
18
+ WS_HOST = os.getenv("FLOWPROXY_WS_HOST", "127.0.0.1")
19
+ EXTENSION_WS_PORT = int(os.getenv("FLOWPROXY_EXT_WS_PORT", "9223"))
20
+
21
+ # Public base used to build absolute media URLs returned to OpenAI clients
22
+ # (response_format="url"). Point this at wherever flow_engine is reachable.
23
+ PUBLIC_BASE_URL = os.getenv("FLOWPROXY_PUBLIC_BASE_URL", f"http://127.0.0.1:{HTTP_PORT}")
24
+
25
+ # Bearer key gating the /v1/* and /api/* endpoints. Prefer the new env var name;
26
+ # fall back to the legacy FLOWPROXY_API_KEY for backwards compatibility.
27
+ # Enforcement (non-empty assertion) happens at server startup in main.py,
28
+ # NOT here, so that `pip install -e .` and module imports don't crash.
29
+ API_KEY = (os.getenv("FLOW_ENGINE_API_KEY") or os.getenv("FLOWPROXY_API_KEY", "")).strip()
30
+
31
+ # Maximum concurrent in-flight Flow API calls per pool instance. Additional
32
+ # callers queue (asyncio.Semaphore semantics — no 503, just back-pressure).
33
+ POOL_MAX_CONCURRENCY = int(os.getenv("FLOW_POOL_CONCURRENCY", "4"))
34
+
35
+ # Title used when flow_engine auto-creates its default Flow project.
36
+ DEFAULT_PROJECT_TITLE = os.getenv("FLOWPROXY_PROJECT_TITLE", "flow_engine")
37
+
38
+ # Model catalog (flow.projectInitialData) is a heavy call — cache it on disk
39
+ # and only refresh once a day. Override the TTL (seconds) or force via
40
+ # /v1/models?refresh=1.
41
+ CATALOG_TTL_S = int(os.getenv("FLOWPROXY_CATALOG_TTL", str(24 * 60 * 60)))
42
+
43
+ STORAGE_DIR.mkdir(parents=True, exist_ok=True)
flow_engine/ingest.py ADDED
@@ -0,0 +1,172 @@
1
+ """Resolve user-supplied reference media into Flow media ids.
2
+
3
+ The generation endpoints (i2i / i2v / r2v) accept reference inputs that are
4
+ EITHER an already-uploaded Flow media id OR raw media the caller wants us to
5
+ upload on their behalf — a ``data:`` URL, an ``http(s)`` URL, or a bare base64
6
+ string. This module classifies each input and, when it's raw, uploads it
7
+ (images via ``uploadImage``, videos via the resumable upload-video flow) so the
8
+ caller ends up with a clean list of media ids ready to dispatch.
9
+
10
+ Already-uploaded ids pass through untouched, so a client can mix pre-uploaded
11
+ references with fresh raw ones in the same request.
12
+ """
13
+ from __future__ import annotations
14
+
15
+ import base64
16
+ import binascii
17
+ import logging
18
+ from typing import Optional
19
+ from urllib.parse import unquote, urlparse
20
+
21
+ import httpx
22
+ from fastapi import HTTPException
23
+
24
+ from flow_engine import media as media_service
25
+ from flow_engine.bridge.flow_sdk import FlowSDK
26
+
27
+ logger = logging.getLogger(__name__)
28
+
29
+ # Cap on bytes pulled from a user-supplied URL — defence against a caller
30
+ # pointing us at an enormous file. 64 MiB comfortably covers an 8s 1080p clip.
31
+ _MAX_FETCH_BYTES = 64 * 1024 * 1024
32
+
33
+ _EXT_BY_MIME = {
34
+ "image/jpeg": ".jpg",
35
+ "image/png": ".png",
36
+ "image/webp": ".webp",
37
+ "image/gif": ".gif",
38
+ "video/mp4": ".mp4",
39
+ "video/webm": ".webm",
40
+ "video/quicktime": ".mov",
41
+ }
42
+
43
+
44
+ async def resolve_image_inputs(
45
+ sdk: FlowSDK, project_id: str, inputs: Optional[list[str]]
46
+ ) -> list[str]:
47
+ """Turn a list of image references into Flow media ids, uploading any raw
48
+ ones. Returns ``[]`` for a falsy/empty input."""
49
+ return await _resolve(sdk, project_id, inputs, want="image")
50
+
51
+
52
+ async def resolve_video_inputs(
53
+ sdk: FlowSDK, project_id: str, inputs: Optional[list[str]]
54
+ ) -> list[str]:
55
+ """Turn a list of video references into Flow media ids, uploading any raw
56
+ ones. Returns ``[]`` for a falsy/empty input."""
57
+ return await _resolve(sdk, project_id, inputs, want="video")
58
+
59
+
60
+ async def _resolve(
61
+ sdk: FlowSDK, project_id: str, inputs: Optional[list[str]], *, want: str
62
+ ) -> list[str]:
63
+ out: list[str] = []
64
+ for raw in inputs or []:
65
+ if not isinstance(raw, str) or not raw.strip():
66
+ continue
67
+ item = raw.strip()
68
+ # Already a Flow media id? Pass it through. A UUID-shaped id is hex +
69
+ # dashes; data:/http: URLs and base64 payloads never match, so this is
70
+ # an unambiguous discriminator.
71
+ candidate = media_service.normalize_media_id(item)
72
+ if media_service.is_valid_media_id(candidate) and not item.startswith(
73
+ ("data:", "http://", "https://")
74
+ ):
75
+ out.append(candidate)
76
+ continue
77
+ b64, mime, name = await _load_raw(item, want=want)
78
+ out.append(await _upload(sdk, project_id, b64, mime, name, want=want))
79
+ return out
80
+
81
+
82
+ async def _upload(
83
+ sdk: FlowSDK,
84
+ project_id: str,
85
+ image_base64: str,
86
+ mime: str,
87
+ file_name: str,
88
+ *,
89
+ want: str,
90
+ ) -> str:
91
+ if want == "video":
92
+ result = await sdk.upload_video(
93
+ video_base64=image_base64,
94
+ mime_type=mime,
95
+ project_id=project_id,
96
+ file_name=file_name,
97
+ )
98
+ else:
99
+ result = await sdk.upload_image(
100
+ image_base64=image_base64,
101
+ mime_type=mime,
102
+ project_id=project_id,
103
+ file_name=file_name,
104
+ )
105
+ media_id = result.get("media_id")
106
+ if result.get("error") or not isinstance(media_id, str) or not media_id:
107
+ detail = str(result.get("error") or "no media_id returned")[:200]
108
+ raise HTTPException(status_code=502, detail=f"reference upload failed: {detail}")
109
+ return media_id
110
+
111
+
112
+ async def _load_raw(item: str, *, want: str) -> tuple[str, str, str]:
113
+ """Return ``(base64_no_prefix, mime, file_name)`` for a raw reference."""
114
+ default_mime = "image/png" if want == "image" else "video/mp4"
115
+ if item.startswith("data:"):
116
+ header, sep, payload = item.partition(",")
117
+ if not sep or not payload:
118
+ raise HTTPException(status_code=400, detail="malformed data URL reference")
119
+ mime = header[len("data:") :].split(";", 1)[0].strip() or default_mime
120
+ _validate_b64(payload)
121
+ return payload, mime, _default_name(mime, want)
122
+ if item.startswith(("http://", "https://")):
123
+ return await _fetch_url(item, want=want)
124
+ # Bare base64 (no envelope) — assume the default mime for the mode.
125
+ _validate_b64(item)
126
+ return item, default_mime, _default_name(default_mime, want)
127
+
128
+
129
+ async def _fetch_url(url: str, *, want: str) -> tuple[str, str, str]:
130
+ try:
131
+ async with httpx.AsyncClient(timeout=30.0, follow_redirects=True) as client:
132
+ resp = await client.get(url)
133
+ except httpx.HTTPError as exc:
134
+ raise HTTPException(status_code=502, detail=f"could not fetch reference url: {exc}")
135
+ if resp.status_code != 200:
136
+ raise HTTPException(
137
+ status_code=502, detail=f"reference url returned HTTP {resp.status_code}"
138
+ )
139
+ content = resp.content
140
+ if len(content) > _MAX_FETCH_BYTES:
141
+ raise HTTPException(status_code=413, detail="reference media exceeds size limit")
142
+ mime = resp.headers.get("content-type", "").split(";")[0].strip().lower()
143
+ prefix = "image/" if want == "image" else "video/"
144
+ if not mime.startswith(prefix):
145
+ raise HTTPException(
146
+ status_code=415,
147
+ detail=f"reference url is not {want} media (content-type {mime or 'unknown'})",
148
+ )
149
+ b64 = base64.b64encode(content).decode("ascii")
150
+ return b64, mime, _name_from_url(url, mime, want)
151
+
152
+
153
+ def _validate_b64(payload: str) -> None:
154
+ """Reject obviously-malformed base64 up front so the upstream upload call
155
+ fails fast with a clear 400 instead of a confusing Flow error."""
156
+ try:
157
+ base64.b64decode(payload, validate=True)
158
+ except (binascii.Error, ValueError):
159
+ raise HTTPException(status_code=400, detail="reference is not valid base64")
160
+
161
+
162
+ def _default_name(mime: str, want: str) -> str:
163
+ ext = _EXT_BY_MIME.get(mime, ".mp4" if want == "video" else ".png")
164
+ return f"upload{ext}"
165
+
166
+
167
+ def _name_from_url(url: str, mime: str, want: str) -> str:
168
+ path = urlparse(url).path
169
+ tail = unquote(path.rsplit("/", 1)[-1]) if path else ""
170
+ if tail and "." in tail:
171
+ return tail
172
+ return _default_name(mime, want)