Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 21 additions & 0 deletions backend/app/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -126,15 +126,30 @@ def setup_logging() -> None:
)


def setup_no_proxy() -> None:
"""Ensure local loopback addresses bypass any HTTP/HTTPS proxies (such as Clash)."""
current_no_proxy = os.environ.get("NO_PROXY", os.environ.get("no_proxy", ""))
needed = ["localhost", "127.0.0.1", "::1", "0.0.0.0"]
existing = [p.strip() for p in current_no_proxy.split(",") if p.strip()]
for n in needed:
if n not in existing:
existing.append(n)
merged = ",".join(existing)
os.environ["NO_PROXY"] = merged
os.environ["no_proxy"] = merged


@asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
setup_logging()
setup_no_proxy()
logger.info("Berry AI Studio API starting up...")
yield
logger.info("Berry AI Studio API shutting down...")


setup_logging()
setup_no_proxy()
logger = logging.getLogger("berry_ai_studio")

app = FastAPI(
Expand Down Expand Up @@ -892,6 +907,12 @@ async def stop_ollama_runtime() -> Dict[str, Any]:
return ollama_supervisor.stop()


@app.post("/api/v1/ollama/install")
async def install_ollama_runtime() -> Dict[str, Any]:
"""Trigger background installation or get installation instructions for Ollama."""
return ollama_supervisor.install()


class OllamaPullModelRequest(BaseModel):
model_name: str

Expand Down
5 changes: 3 additions & 2 deletions backend/app/runners/comfy_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,9 +38,10 @@ def __init__(self, host: str = "127.0.0.1", port: int = 8188, timeout: float = 5
self._client: Optional[httpx.AsyncClient] = None

def _get_client(self) -> httpx.AsyncClient:
"""Reuse long-lived pooled client to prevent socket exhaustion."""
"""Reuse long-lived pooled client to prevent socket exhaustion, bypassing proxy for local host."""
if self._client is None or self._client.is_closed:
self._client = httpx.AsyncClient(timeout=self.timeout)
is_local = self.host in ("127.0.0.1", "localhost", "0.0.0.0")
self._client = httpx.AsyncClient(timeout=self.timeout, trust_env=not is_local)
return self._client

async def aclose(self) -> None:
Expand Down
6 changes: 4 additions & 2 deletions backend/app/runners/webui_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,8 @@ def __init__(self, endpoint_url: str = "http://127.0.0.1:7860") -> None:

async def execute_action(self, req: CreativeActionRequest) -> Dict[str, Any]:
"""Execute creative action via SD WebUI REST API."""
async with httpx.AsyncClient(timeout=120.0) as client:
is_local = "127.0.0.1" in self.endpoint_url or "localhost" in self.endpoint_url
async with httpx.AsyncClient(timeout=120.0, trust_env=not is_local) as client:
if req.action == CreativeActionType.TXT2IMG:
return await self._run_txt2img(client, req)
elif req.action == CreativeActionType.IMG2IMG:
Expand All @@ -39,7 +40,8 @@ async def execute_action(self, req: CreativeActionRequest) -> Dict[str, Any]:
async def interrupt(self) -> bool:
"""Interrupt active execution on SD WebUI via POST /sdapi/v1/interrupt."""
try:
async with httpx.AsyncClient(timeout=5.0) as client:
is_local = "127.0.0.1" in self.endpoint_url or "localhost" in self.endpoint_url
async with httpx.AsyncClient(timeout=5.0, trust_env=not is_local) as client:
resp = await client.post(f"{self.endpoint_url}/sdapi/v1/interrupt")
return resp.status_code == 200
except Exception as err:
Expand Down
3 changes: 2 additions & 1 deletion backend/app/runtime/engine_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,7 +116,8 @@ async def connect_external_engine(
async def test_engine_connection(self, connection: EngineConnection) -> EngineConnection:
"""Probe engine endpoint and update status and capabilities."""
endpoint = connection.endpoint_url
async with httpx.AsyncClient(timeout=4.0) as client:
is_local = "127.0.0.1" in endpoint or "localhost" in endpoint
async with httpx.AsyncClient(timeout=4.0, trust_env=not is_local) as client:
try:
if connection.engine_type == EngineType.COMFYUI:
# Query /system_stats
Expand Down
90 changes: 84 additions & 6 deletions backend/app/runtime/ollama_supervisor.py
Original file line number Diff line number Diff line change
Expand Up @@ -119,9 +119,9 @@ def is_running(self) -> bool:
return self.get_pid() is not None

async def check_health(self) -> bool:
"""Ping local Ollama HTTP endpoint."""
"""Ping local Ollama HTTP endpoint bypassing any system proxies."""
try:
async with httpx.AsyncClient(timeout=2.0) as client:
async with httpx.AsyncClient(timeout=2.0, trust_env=False) as client:
res = await client.get(f"http://127.0.0.1:{self.port}/api/tags")
return res.status_code == 200
except Exception:
Expand All @@ -139,7 +139,7 @@ async def get_status(self) -> OllamaRuntimeStatus:
if healthy:
running = True
try:
async with httpx.AsyncClient(timeout=3.0) as client:
async with httpx.AsyncClient(timeout=3.0, trust_env=False) as client:
tags_res = await client.get(f"http://127.0.0.1:{self.port}/api/tags")
if tags_res.status_code == 200:
data = tags_res.json()
Expand Down Expand Up @@ -233,16 +233,94 @@ def stop(self) -> Dict[str, Any]:

return {"success": True, "message": f"Ollama process {pid} stopped"}

def install(self) -> Dict[str, Any]:
"""Attempt to install Ollama runtime on the host system."""
if self.is_installed():
return {
"success": True,
"message": "Ollama is already installed",
"installed": True,
}

if sys.platform == "win32":
winget_cmd = shutil.which("winget")
if winget_cmd:
try:
subprocess.Popen(
[winget_cmd, "install", "Ollama.Ollama", "--accept-source-agreements", "--accept-package-agreements", "--silent"],
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
creationflags=subprocess.CREATE_NO_WINDOW,
)
return {
"success": True,
"message": "Installing Ollama in background via winget. Please wait 1-2 minutes.",
"installing": True,
"method": "winget",
}
except Exception as e:
logger.warning(f"winget installation failed: {e}")

return {
"success": False,
"message": "Please install Ollama from https://ollama.com",
"installing": False,
"download_url": "https://ollama.com/download/windows",
}
else:
return {
"success": False,
"message": "Please install Ollama via: curl -fsSL https://ollama.com/install.sh | sh",
"installing": False,
"download_url": "https://ollama.com",
}

async def pull_model_stream(self, model_name: str) -> AsyncGenerator[Dict[str, Any], None]:
"""Stream model pulling progress from Ollama /api/pull."""
"""Stream model pulling progress from Ollama /api/pull, ensuring Ollama is active and bypassing proxies."""
if not self.is_installed():
yield {
"status": "error",
"error": "Ollama is not installed. Please install Ollama from https://ollama.com first.",
"code": "NOT_INSTALLED",
}
return

# Ensure Ollama daemon is running, auto-start if needed
if not (await self.check_health()):
logger.info("Ollama is not running. Auto-starting Ollama service...")
start_res = self.start()
if not start_res.get("success") and not self.is_running():
yield {
"status": "error",
"error": f"Failed to start local Ollama engine: {start_res.get('message', 'Unknown error')}",
"code": "START_FAILED",
}
return

# Wait up to 6 seconds for Ollama HTTP endpoint to become healthy
ready = False
for _ in range(12):
await asyncio.sleep(0.5)
if await self.check_health():
ready = True
break

if not ready:
yield {
"status": "error",
"error": f"Ollama service started but is not responding on port {self.port}.",
"code": "PORT_UNRESPONSIVE",
}
return

url = f"http://127.0.0.1:{self.port}/api/pull"
self._pulling_tasks[model_name] = {"status": "pulling", "total": 0, "completed": 0}

try:
async with httpx.AsyncClient(timeout=None) as client:
async with httpx.AsyncClient(timeout=None, trust_env=False) as client:
async with client.stream("POST", url, json={"name": model_name, "stream": True}) as response:
if response.status_code != 200:
yield {"status": "error", "error": f"Failed to pull model ({response.status_code})"}
yield {"status": "error", "error": f"Ollama returned HTTP error ({response.status_code})"}
return

async for line in response.aiter_lines():
Expand Down
54 changes: 54 additions & 0 deletions backend/tests/test_ollama_supervisor.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
import os
import pytest
from unittest.mock import patch, AsyncMock, MagicMock
from fastapi.testclient import TestClient
from app.main import app, setup_no_proxy
from app.runtime.ollama_supervisor import ollama_supervisor, OllamaSupervisor


def test_setup_no_proxy():
setup_no_proxy()
no_proxy = os.environ.get("NO_PROXY", "")
assert "127.0.0.1" in no_proxy
assert "localhost" in no_proxy


@pytest.mark.asyncio
async def test_pull_model_stream_not_installed():
supervisor = OllamaSupervisor()
with patch.object(supervisor, "is_installed", return_value=False):
events = []
async for chunk in supervisor.pull_model_stream("qwen2.5:14b"):
events.append(chunk)

assert len(events) == 1
assert events[0]["status"] == "error"
assert "not installed" in events[0]["error"].lower()
assert events[0].get("code") == "NOT_INSTALLED"


@pytest.mark.asyncio
async def test_pull_model_stream_auto_start_failure():
supervisor = OllamaSupervisor()
with patch.object(supervisor, "is_installed", return_value=True), \
patch.object(supervisor, "check_health", new_callable=AsyncMock, return_value=False), \
patch.object(supervisor, "start", return_value={"success": False, "message": "Exec error"}), \
patch.object(supervisor, "is_running", return_value=False):

events = []
async for chunk in supervisor.pull_model_stream("qwen2.5:7b"):
events.append(chunk)

assert len(events) == 1
assert events[0]["status"] == "error"
assert "failed to start" in events[0]["error"].lower()


def test_ollama_install_endpoint():
client = TestClient(app)
with patch.object(ollama_supervisor, "install", return_value={"success": True, "message": "Installing", "installing": True}):
res = client.post("/api/v1/ollama/install")
assert res.status_code == 200
data = res.json()
assert data["success"] is True
assert data["installing"] is True
Loading
Loading