diff --git a/ai_diffusion/backend/comfy_client.py b/ai_diffusion/backend/comfy_client.py index 95d292934..724077695 100644 --- a/ai_diffusion/backend/comfy_client.py +++ b/ai_diffusion/backend/comfy_client.py @@ -269,14 +269,17 @@ async def refresh(self): async for __ in self.discover_models(refresh=True): pass + async def _head(self, op: str, timeout: float | None = 60): + return await self._requests.head(f"{self.url}/{op}", timeout=timeout) + async def _get(self, op: str, timeout: float | None = 60): return await self._requests.get(f"{self.url}/{op}", timeout=timeout) async def _post(self, op: str, data: dict): return await self._requests.post(f"{self.url}/{op}", data) - async def _put(self, op: str, data: bytes): - return await self._requests.put(f"{self.url}/{op}", data) + async def _put(self, op: str, data: bytes, expect_continue: bool = False): + return await self._requests.put(f"{self.url}/{op}", data, expect_continue=expect_continue) async def enqueue(self, work: WorkflowInput, front: bool = False): job = JobInfo.create(work) @@ -548,7 +551,7 @@ async def _transfer_result_image(self, id: str): async def upload_images(self, image_data: dict[str, bytes]): for id, data in image_data.items(): try: - await self._put(f"api/etn/image/{id}", data) + await self._put(f"api/etn/image/{id}", data, expect_continue=True) except Exception as e: log.error(f"Error uploading image {id}: {e!s}") raise RuntimeError(f"Error uploading input image to ComfyUI: {e!s}") from e diff --git a/ai_diffusion/backend/network.py b/ai_diffusion/backend/network.py index f106392ec..b66997521 100644 --- a/ai_diffusion/backend/network.py +++ b/ai_diffusion/backend/network.py @@ -9,7 +9,14 @@ from typing import NamedTuple from PyQt5.QtCore import QBuffer, QByteArray, QFile, QUrl -from PyQt5.QtNetwork import QNetworkAccessManager, QNetworkReply, QNetworkRequest, QSslError +from PyQt5.QtNetwork import ( + QNetworkAccessManager, + QNetworkReply, + QNetworkRequest, + QSslError, + QSslSocket, + QTcpSocket, +) from ..localization import translate as _ from ..util import client_logger as log @@ -77,6 +84,313 @@ class Request(NamedTuple): Headers = list[tuple[str, str]] +class _SocketPutRequest: + def __init__( + self, + manager: "RequestManager", + url: str, + data: QByteArray, + timeout: float | None, + bearer: str | None, + expect_continue: bool, + ): + self._manager = manager + self._loop = asyncio.get_running_loop() + self._url = QUrl(url) + self._is_https = self._url.scheme().lower() == "https" + self._data = bytes(data) + self._timeout = timeout + self._bearer = bearer + self._expect_continue = expect_continue + self._future = self._loop.create_future() + + self._socket = QSslSocket() if self._is_https else QTcpSocket() + self._buffer = bytearray() + self._headers_sent = False + self._waiting_for_continue = False + self._final_status: int | None = None + self._final_headers: dict[str, str] = {} + self._final_body = bytearray() + self._final_content_length: int | None = None + self._upload_started = False + self._upload_offset = 0 + + self._connect_signals() + self._connect_socket() + + @property + def future(self) -> asyncio.Future: + return self._future + + def _connect_signals(self): + if self._is_https and isinstance(self._socket, QSslSocket): + self._socket.encrypted.connect(self._on_connected) + self._socket.sslErrors.connect(self._on_ssl_errors) + else: + self._socket.connected.connect(self._on_connected) + self._socket.readyRead.connect(self._on_ready_read) + self._socket.bytesWritten.connect(self._on_bytes_written) + self._socket.errorOccurred.connect(self._on_socket_error) + self._socket.disconnected.connect(self._on_disconnected) + + def _connect_socket(self): + host = self._url.host() + default_port = 443 if self._is_https else 80 + port = self._url.port(default_port) + if self._is_https and isinstance(self._socket, QSslSocket): + log.debug(f"Connecting securely to {host}:{port} for PUT request") + self._socket.connectToHostEncrypted(host, port) + else: + log.debug(f"Connecting to {host}:{port} for PUT request") + self._socket.connectToHost(host, port) + + def _on_connected(self): + log.debug(f"Connected to {self._url.host()} for PUT request") + if self._future.done(): + return + + log.debug(f"Sending HTTP PUT request to {self._url.toString()}") + path = self._url.path() or "/" + if self._url.hasQuery(): + path += f"?{self._url.query()}" + + host = self._url.host() + default_port = 443 if self._is_https else 80 + port = self._url.port(default_port) + host_header = f"{host}:{port}" if port != default_port else host + + headers: list[tuple[str, str]] = [ + ("Host", host_header), + ("User-Agent", "krita-ai-diffusion"), + ("Accept", "*/*"), + ("Content-Type", "application/octet-stream"), + ("Content-Length", str(len(self._data))), + ] + if self._expect_continue: + log.debug("Adding Expect: 100-continue header for PUT request") + headers.append(("Expect", "100-continue")) + + headers.extend(self._manager.socket_headers(self._bearer)) + + request = [f"PUT {path} HTTP/1.1"] + request.extend([f"{k}: {v}" for k, v in headers]) + request_bytes = "\r\n".join(request).encode("utf-8") + b"\r\n\r\n" + + self._socket.write(request_bytes) + self._headers_sent = True + self._waiting_for_continue = self._expect_continue + + if self._expect_continue: + # Some servers ignore Expect: 100-continue and expect the payload immediately. + # Wait long enough for a round-trip even on slow internet/cloud connections. + self._loop.call_later(2.0, self._ensure_upload_started) + else: + self._start_upload() + if self._timeout is not None: + self._loop.call_later(self._timeout, self._on_timeout) + + def _on_timeout(self): + if self._future.done(): + return + self._socket.abort() + self._fail( + NetworkError( + int(QNetworkReply.NetworkError.TimeoutError), + "Connection timed out, the server took too long to respond", + self._url.toString(), + ) + ) + + def _ensure_upload_started(self): + if self._future.done() or self._upload_started or self._final_status is not None: + return + self._start_upload() + + def _start_upload(self): + if self._future.done() or self._upload_started or self._final_status is not None: + return + self._upload_started = True + self._waiting_for_continue = False + self._pump_upload() + + def _on_bytes_written(self, _count: int): + if self._future.done() or not self._upload_started: + return + self._pump_upload() + + def _pump_upload(self): + if self._future.done() or self._final_status is not None: + return + + # Keep outgoing buffer bounded so incoming data can be processed quickly. + max_inflight = 256 * 1024 + chunk_size = 64 * 1024 + + log.debug(f"Upload pump for {self._url.toString()}: offset {self._upload_offset}/{len(self._data)} bytes, bytesToWrite={self._socket.bytesToWrite()}") + while self._socket.bytesToWrite() < max_inflight and self._upload_offset < len(self._data): + end = min(self._upload_offset + chunk_size, len(self._data)) + chunk = self._data[self._upload_offset : end] + log.debug(f"Uploading chunk of size {len(chunk)} bytes at offset {self._upload_offset}/{len(self._data)} bytes.") + self._upload_offset = end + self._socket.write(chunk) + log.debug(f"Chunk upload complete, offset {self._upload_offset}/{len(self._data)} bytes, bytesToWrite={self._socket.bytesToWrite()}") + log.debug(f"Upload pump for {self._url.toString()} complete, offset {self._upload_offset}/{len(self._data)} bytes, bytesToWrite={self._socket.bytesToWrite()}") + + def _on_ready_read(self): + if self._future.done(): + return + + self._buffer.extend(bytes(self._socket.readAll())) + try: + self._process_http_stream() + except Exception as e: + self._fail(e) + + def _process_http_stream(self): + log.debug(f"Processing HTTP stream from {self._url.toString()}, buffer size {len(self._buffer)} bytes") + while True: + if self._final_status is None: + parsed = self._try_parse_response_headers() + if parsed is None: + return + + status, headers = parsed + log.debug(f"Received HTTP response with status {status} from {self._url.toString()}") + if status == 100: + log.debug(f"Received 100 Continue from {self._url.toString()}") + self._waiting_for_continue = False + self._start_upload() + continue + + self._final_status = status + self._final_headers = headers + content_length = headers.get("content-length") + if content_length is not None: + try: + self._final_content_length = int(content_length) + except ValueError: + self._final_content_length = None + + if self._final_content_length is None: + log.debug(f"No Content-Length for response from {self._url.toString()}, buffering until disconnect") + # Body length unknown; wait for disconnect to complete. + self._final_body.extend(self._buffer) + self._buffer.clear() + return + + if len(self._buffer) < self._final_content_length: + return + + self._final_body.extend(self._buffer[: self._final_content_length]) + del self._buffer[: self._final_content_length] + self._complete_from_final_response() + return + + def _try_parse_response_headers(self) -> tuple[int, dict[str, str]] | None: + marker = self._buffer.find(b"\r\n\r\n") + if marker < 0: + return None + + raw_header = bytes(self._buffer[:marker]) + del self._buffer[: marker + 4] + lines = raw_header.decode("latin-1", errors="replace").split("\r\n") + if not lines or not lines[0].startswith("HTTP/"): + raise RuntimeError(f"Invalid HTTP response from {self._url.toString()}") + + status_parts = lines[0].split(" ", 2) + try: + status = int(status_parts[1]) + except Exception as e: + raise RuntimeError(f"Invalid HTTP status from {self._url.toString()}") from e + + headers: dict[str, str] = {} + for line in lines[1:]: + if not line or ":" not in line: + continue + key, value = line.split(":", 1) + headers[key.strip().lower()] = value.strip() + return status, headers + + def _on_disconnected(self): + if self._future.done(): + return + + if self._final_status is not None: + self._final_body.extend(self._buffer) + self._buffer.clear() + self._complete_from_final_response() + return + + self._fail( + NetworkError( + int(QNetworkReply.NetworkError.RemoteHostClosedError), + "Connection closed before response was received", + self._url.toString(), + ) + ) + + def _on_socket_error(self, _error: QTcpSocket.SocketError): + if self._future.done(): + return + + # If final HTTP response was already seen, a disconnect error is expected. + if self._final_status is not None: + return + + self._fail( + NetworkError( + int(QNetworkReply.NetworkError.NetworkSessionFailedError), + self._socket.errorString(), + self._url.toString(), + ) + ) + + def _on_ssl_errors(self, errors: list[QSslError]): + if self._future.done(): + return + error_text = "; ".join(error.errorString() for error in errors) or "SSL handshake failed" + self._fail( + NetworkError( + int(QNetworkReply.NetworkError.SslHandshakeFailedError), + error_text, + self._url.toString(), + ) + ) + + def _complete_from_final_response(self): + if self._future.done() or self._final_status is None: + log.warning(f"Attempted to complete request for {self._url.toString()} but final status is not set") + return + + status = self._final_status + data = bytes(self._final_body) + log.debug(f"Final response from {self._url.toString()}: status {status}, body size {len(data)} bytes") + if 200 <= status < 300: + content_type = self._final_headers.get("content-type", "") + if "application/json" in content_type: + self._future.set_result(json.loads(data or b"{}")) + log.debug(f"Parsed JSON response from {self._url.toString()}: {self._future.result()}") + else: + self._future.set_result(data) + self._socket.abort() + self._manager.release_socket_put(self) + return + + msg = data.decode("utf-8", errors="replace") if data else f"HTTP {status}" + self._fail( + NetworkError( + int(QNetworkReply.NetworkError.UnknownServerError), + msg, + self._url.toString(), + status, + ) + ) + + def _fail(self, error: Exception): + if not self._future.done(): + self._future.set_exception(error) + self._socket.abort() + self._manager.release_socket_put(self) class RequestManager: def __init__(self): @@ -85,6 +399,7 @@ def __init__(self): self._net.sslErrors.connect(self._handle_ssl_errors) self._requests: dict[QNetworkReply, Request] = {} self._upload_future: Future[tuple[int, int]] | None = None + self._socket_puts: set[_SocketPutRequest] = set() self._additional_headers: list[tuple[bytes, bytes]] = [] self._bearer_token: str | None = None @@ -106,6 +421,18 @@ def _prepare_request(self, url: str, timeout: float | None = None, bearer: str | request.setTransferTimeout(int(timeout * 1000)) return request + def socket_headers(self, bearer: str | None = None): + headers: list[tuple[str, str]] = [] + bearer_token = bearer or self._bearer_token + if bearer_token: + headers.append(("Authorization", f"Bearer {bearer_token}")) + for key, value in self._additional_headers: + headers.append((key.decode("utf-8"), value.decode("utf-8"))) + return headers + + def release_socket_put(self, request: _SocketPutRequest): + self._socket_puts.discard(request) + def http( self, method, @@ -118,7 +445,7 @@ def http( request = self._prepare_request(url, timeout, bearer) - assert method in ["GET", "POST", "PUT"] + assert method in ["GET", "POST", "PUT", "HEAD"] if method == "POST": data = data or {} data_bytes = QByteArray(json.dumps(data).encode("utf-8")) @@ -134,6 +461,8 @@ def http( ) request.setHeader(QNetworkRequest.KnownHeaders.ContentLengthHeader, data.size()) reply = self._net.put(request, data) + elif method == "HEAD": + reply = self._net.head(request) else: reply = self._net.get(request) @@ -148,8 +477,48 @@ def get(self, url: str, timeout: float | None = None, bearer: str | None = None) def post(self, url: str, data: dict, bearer: str | None = None): return self.http("POST", url, data, bearer=bearer) - def put(self, url: str, data: QByteArray | bytes): - return self.http("PUT", url, data) + def put( + self, + url: str, + data: QByteArray | bytes, + timeout: float | None = None, + use_socket: bool = False, + expect_continue: bool = False, + bearer: str | None = None, + ): + # When Expect: 100-continue is requested, prefer socket transport so early + # final responses can be handled before full payload upload. + if use_socket or expect_continue: + return self.put_socket( + url, + data, + timeout=timeout, + bearer=bearer, + expect_continue=expect_continue, + ) + if not isinstance(data, QByteArray): + data = QByteArray(bytes(data)) + return self.http("PUT", url, data, timeout=timeout, bearer=bearer) + + def put_socket( + self, + url: str, + data: QByteArray | bytes, + timeout: float | None = None, + bearer: str | None = None, + expect_continue: bool = True, + ): + self._cleanup() + if not isinstance(data, QByteArray): + data = QByteArray(bytes(data)) + assert isinstance(data, QByteArray) + + request = _SocketPutRequest(self, url, data, timeout, bearer, expect_continue) + self._socket_puts.add(request) + return request.future + + def head(self, url: str, timeout: float | None = None, bearer: str | None = None): + return self.http("HEAD", url, timeout=timeout, bearer=bearer) async def upload(self, url: str, data: QByteArray | bytes, sha256: str | None = None): self._cleanup() diff --git a/tests/test_network.py b/tests/test_network.py new file mode 100644 index 000000000..760f33361 --- /dev/null +++ b/tests/test_network.py @@ -0,0 +1,323 @@ +"""Tests for network module, especially socket-based PUT with Expect: 100-continue.""" + +from __future__ import annotations + +import asyncio +import json +import socket + +import pytest +from aiohttp import web +from PyQt5.QtCore import QByteArray + +from ai_diffusion.backend.network import NetworkError, RequestManager + +from .conftest import qtapp + + +# --------------------------------------------------------------------------- +# Test HTTP Server +# --------------------------------------------------------------------------- + + +class _TestHTTPServer: + """Async test HTTP server using aiohttp for testing socket PUT requests.""" + + def __init__(self, port: int | None = None): + self.port = port or _find_free_port() + self.app = web.Application(client_max_size=16 * 1024**2) + self.runner: web.AppRunner | None = None + self.site: web.TCPSite | None = None + self.behavior = "send_100_continue" + self.request_body: bytearray | None = None + + # Route PUT requests to the handler + self.app.router.add_put("/api/upload", self._handle_put) + + async def _handle_put(self, request: web.Request) -> web.Response: + """Handle PUT requests according to configured behavior.""" + if self.behavior == "send_100_continue": + # Read payload and send 200 OK (normal upload) + body = await request.read() + self.request_body = bytearray(body) + response = {"status": "ok"} + return web.json_response(response, status=200) + + elif self.behavior == "send_early_200": + # Send 200 OK immediately without reading payload (matches production expect_handler) + # In production, this uses aiohttp's expect_handler to detect cached images before + # body consumption. The client sees 200 and skips the upload. + self.request_body = None + response = {"status": "cached"} + return web.json_response(response, status=200) + + elif self.behavior == "ignore_expect": + # Read payload normally, send 200 OK (server ignores Expect: 100-continue) + body = await request.read() + self.request_body = bytearray(body) + response = {"status": "ok"} + return web.json_response(response, status=200) + + elif self.behavior == "send_error": + # Send error response + response = {"error": "Test error"} + return web.json_response(response, status=400) + + elif self.behavior == "slow_100_continue": + # Simulate slow server + await asyncio.sleep(0.5) + body = await request.read() + self.request_body = bytearray(body) + response = {"status": "ok"} + return web.json_response(response, status=200) + + else: + return web.json_response({"error": "Unknown behavior"}, status=500) + + async def start(self): + """Start the server.""" + self.runner = web.AppRunner(self.app) + await self.runner.setup() + self.site = web.TCPSite(self.runner, "127.0.0.1", self.port) + await self.site.start() + + async def stop(self): + """Stop the server.""" + if self.runner is not None: + await self.runner.cleanup() + + def set_behavior(self, behavior: str): + """Set how the server should respond to requests.""" + self.behavior = behavior + self.request_body = None + + @property + def url(self) -> str: + return f"http://127.0.0.1:{self.port}" + + +def _find_free_port(): + """Find an available port on localhost.""" + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("127.0.0.1", 0)) + s.listen(1) + port = s.getsockname()[1] + return port + + +@pytest.fixture +def test_server(qtapp): + """Fixture providing an aiohttp test server.""" + server = _TestHTTPServer() + qtapp.run(server.start()) + yield server + qtapp.run(server.stop()) + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + + +@qtapp +async def test_socket_put_with_100_continue(test_server: _TestHTTPServer): + """Socket PUT receives 100 Continue, then uploads payload, then receives 200 OK.""" + test_server.set_behavior("send_100_continue") + + manager = RequestManager() + test_data = b"x" * (100 * 1024) # 100 KB + + result = await manager.put_socket( + f"{test_server.url}/api/upload", + test_data, + timeout=5.0, + expect_continue=True, + ) + + assert isinstance(result, dict) + assert result.get("status") == "ok" + assert test_server.request_body == bytearray(test_data) + + +@qtapp +async def test_socket_put_early_200_skips_upload(test_server: _TestHTTPServer): + """Socket PUT receives early 200 OK (cached); upload is skipped entirely.""" + test_server.set_behavior("send_early_200") + + manager = RequestManager() + test_data = b"x" * (100 * 1024) # 100 KB + + result = await manager.put_socket( + f"{test_server.url}/api/upload", + test_data, + timeout=5.0, + expect_continue=True, + ) + + assert isinstance(result, dict) + assert result.get("status") == "cached" + # Payload should NOT be sent to the server in this case + assert test_server.request_body is None or len(test_server.request_body) == 0 + + +@qtapp +async def test_socket_put_ignore_expect_fallback(test_server: _TestHTTPServer): + """Socket PUT with server ignoring Expect: fallback timer triggers upload after 2s.""" + test_server.set_behavior("ignore_expect") + + manager = RequestManager() + test_data = b"y" * (100 * 1024) # 100 KB + + result = await manager.put_socket( + f"{test_server.url}/api/upload", + test_data, + timeout=5.0, + expect_continue=True, + ) + + assert isinstance(result, dict) + assert result.get("status") == "ok" + assert test_server.request_body == bytearray(test_data) + + +@qtapp +async def test_socket_put_error_response(test_server: _TestHTTPServer): + """Socket PUT receives error status code (400).""" + test_server.set_behavior("send_error") + + manager = RequestManager() + test_data = b"z" * (10 * 1024) + + with pytest.raises(NetworkError) as exc_info: + await manager.put_socket( + f"{test_server.url}/api/upload", + test_data, + timeout=5.0, + expect_continue=True, + ) + + error = exc_info.value + assert error.status == 400 + assert "error" in error.message.lower() or "Test error" in error.message + + +@qtapp +async def test_socket_put_without_expect_continue(test_server: _TestHTTPServer): + """Socket PUT without Expect: 100-continue header; immediate upload.""" + test_server.set_behavior("ignore_expect") + + manager = RequestManager() + test_data = b"w" * (50 * 1024) + + result = await manager.put_socket( + f"{test_server.url}/api/upload", + test_data, + timeout=5.0, + expect_continue=False, + ) + + assert isinstance(result, dict) + assert result.get("status") == "ok" + assert test_server.request_body == bytearray(test_data) + + +@qtapp +async def test_socket_put_timeout_on_no_response(test_server: _TestHTTPServer): + """Socket PUT times out if server never responds.""" + # Don't set any behavior; server won't respond at all + # Instead, we'll test against a port with no server + manager = RequestManager() + test_data = b"x" * (10 * 1024) + + with pytest.raises(NetworkError) as exc_info: + await manager.put_socket( + "http://127.0.0.1:1", # Port 1 is reserved; connection will fail + test_data, + timeout=1.0, + expect_continue=True, + ) + + # Should fail due to connection error or timeout + error = exc_info.value + assert error.code is not None + + +@qtapp +async def test_socket_put_small_payload(test_server: _TestHTTPServer): + """Socket PUT with very small payload.""" + test_server.set_behavior("send_100_continue") + + manager = RequestManager() + test_data = b"hello" + + result = await manager.put_socket( + f"{test_server.url}/api/upload", + test_data, + timeout=5.0, + expect_continue=True, + ) + + assert isinstance(result, dict) + assert result.get("status") == "ok" + assert test_server.request_body == bytearray(test_data) + + +@qtapp +async def test_socket_put_large_payload(test_server: _TestHTTPServer): + """Socket PUT with large payload (multiple chunks).""" + test_server.set_behavior("send_100_continue") + + manager = RequestManager() + test_data = b"z" * (5 * 1024 * 1024) # 5 MB + + result = await manager.put_socket( + f"{test_server.url}/api/upload", + test_data, + timeout=10.0, + expect_continue=True, + ) + + assert isinstance(result, dict) + assert result.get("status") == "ok" + assert test_server.request_body == bytearray(test_data) + + +@qtapp +async def test_socket_put_qbytearray_input(test_server: _TestHTTPServer): + """Socket PUT accepts QByteArray as input.""" + test_server.set_behavior("send_100_continue") + + manager = RequestManager() + test_data = b"qt_test_data" + q_data = QByteArray(test_data) + + result = await manager.put_socket( + f"{test_server.url}/api/upload", + q_data, + timeout=5.0, + expect_continue=True, + ) + + assert isinstance(result, dict) + assert result.get("status") == "ok" + assert test_server.request_body == bytearray(test_data) + + +@qtapp +async def test_socket_put_with_bearer_token(test_server: _TestHTTPServer): + """Socket PUT includes Authorization header when bearer token is set.""" + test_server.set_behavior("send_100_continue") + + manager = RequestManager() + manager.set_auth("test-token-12345") + test_data = b"data" + + result = await manager.put_socket( + f"{test_server.url}/api/upload", + test_data, + timeout=5.0, + expect_continue=True, + ) + + assert isinstance(result, dict) + assert result.get("status") == "ok"