Fix coordinator worker integration

This commit is contained in:
Emil
2026-07-23 21:33:37 +03:00
parent b4a89dd7c2
commit 983c5843ec
25 changed files with 748 additions and 158 deletions
+32 -41
View File
@@ -7,38 +7,23 @@ import http.client
import json
from pathlib import Path
from typing import Protocol
from urllib.parse import quote, urlsplit
from urllib.request import HTTPRedirectHandler, Request, build_opener
from urllib.parse import quote, urljoin, urlsplit
from urllib.request import Request, build_opener
from .models import ClaimedTask, ProducedArtifact
from .coordinator import CoordinatorConflictError
from .models import ClaimedTask, ProducedArtifact, UploadedArtifact
from .transport import SameOriginAuthRedirectHandler, origin
def _origin(uri: str) -> tuple[str, str, int | None]:
parsed = urlsplit(uri)
scheme = parsed.scheme.lower()
default_port = {"http": 80, "https": 443}.get(scheme)
return scheme, (parsed.hostname or "").lower(), parsed.port or default_port
class _SameOriginAuthRedirectHandler(HTTPRedirectHandler):
"""Do not forward the coordinator token when a download changes origin."""
def __init__(self, coordinator_origin: tuple[str, str, int | None]) -> None:
super().__init__()
self.coordinator_origin = coordinator_origin
def redirect_request(self, req: Request, fp: object, code: int, msg: str, headers: object, newurl: str) -> Request | None:
redirected = super().redirect_request(req, fp, code, msg, headers, newurl)
if redirected and _origin(newurl) != self.coordinator_origin:
redirected.remove_header("Authorization")
return redirected
# Compatibility aliases for focused transport tests.
_SameOriginAuthRedirectHandler = SameOriginAuthRedirectHandler
_origin = origin
class ArtifactClient(Protocol):
def download(self, uri: str, destination: Path) -> None: ...
def upload(
self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact
) -> str: ...
) -> UploadedArtifact: ...
class HttpArtifactClient:
@@ -48,18 +33,21 @@ class HttpArtifactClient:
self.coordinator_url = coordinator_url.rstrip("/")
self.timeout = timeout
self.bearer_token = bearer_token
self.coordinator_origin = _origin(coordinator_url)
self._opener = build_opener(_SameOriginAuthRedirectHandler(self.coordinator_origin))
self.coordinator_origin = origin(coordinator_url)
self._opener = build_opener(SameOriginAuthRedirectHandler(self.coordinator_origin))
def download(self, uri: str, destination: Path) -> None:
destination.parent.mkdir(parents=True, exist_ok=True)
request = Request(uri, headers=self._auth_headers_for(uri))
resolved_uri = urljoin(f"{self.coordinator_url}/", uri)
request = Request(resolved_uri, headers=self._auth_headers_for(resolved_uri))
with self._opener.open(request, timeout=self.timeout) as response, destination.open("wb") as target:
while chunk := response.read(1024 * 1024):
target.write(chunk)
def upload(self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact) -> str:
"""Stream one result artifact to the coordinator and return its stable URI."""
def upload(
self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact
) -> UploadedArtifact:
"""Stream an artifact and require durable coordinator-owned metadata."""
url = (
f"{self.coordinator_url}/tasks/{quote(task.task_id, safe='')}/artifacts/"
f"{quote(artifact.path.name, safe='')}"
@@ -71,11 +59,13 @@ class HttpArtifactClient:
http.client.HTTPSConnection if parsed.scheme == "https" else http.client.HTTPConnection
)
connection = connection_class(parsed.hostname, parsed.port, timeout=self.timeout)
local_size = artifact.path.stat().st_size
local_sha256 = sha256_file(artifact.path)
try:
path = parsed.path + (f"?{parsed.query}" if parsed.query else "")
connection.putrequest("PUT", path)
connection.putheader("Content-Type", artifact.content_type)
connection.putheader("Content-Length", str(artifact.path.stat().st_size))
connection.putheader("Content-Length", str(local_size))
connection.putheader("X-Worker-ID", worker_id)
connection.putheader("X-Task-Attempt", str(task.attempt))
for name, value in self._auth_headers_for(url).items():
@@ -86,23 +76,24 @@ class HttpArtifactClient:
connection.send(chunk)
response = connection.getresponse()
body = response.read()
if not 200 <= response.status < 300:
if response.status == 409:
raise CoordinatorConflictError("artifact upload rejected because the task lease was lost")
if response.status != 200:
raise RuntimeError(f"artifact upload rejected with status {response.status}")
if body:
try:
response_data = json.loads(body)
except json.JSONDecodeError as error:
raise RuntimeError("artifact upload returned invalid JSON") from error
response_uri = response_data.get("uri") if isinstance(response_data, dict) else None
if isinstance(response_uri, str) and response_uri:
return response_uri
return url
try:
response_data = json.loads(body)
uploaded = UploadedArtifact.from_json(response_data)
except (ValueError, json.JSONDecodeError) as error:
raise RuntimeError("artifact upload returned invalid metadata") from error
if uploaded.sha256 != local_sha256 or uploaded.size_bytes != local_size:
raise RuntimeError("artifact upload metadata does not match local artifact")
return uploaded
finally:
connection.close()
def _auth_headers_for(self, uri: str) -> dict[str, str]:
"""Only coordinator-owned URLs receive the coordinator bearer token."""
if self.bearer_token and _origin(uri) == self.coordinator_origin:
if self.bearer_token and origin(uri) == self.coordinator_origin:
return {"Authorization": f"Bearer {self.bearer_token}"}
return {}