diff --git a/.changeset/cancel-active-direct-highs.md b/.changeset/cancel-active-direct-highs.md new file mode 100644 index 000000000..3190e7be6 --- /dev/null +++ b/.changeset/cancel-active-direct-highs.md @@ -0,0 +1,6 @@ +--- +"ftw": patch +--- + +Cancel no-longer-needed direct-HiGHS solves in the Unix optimizer sidecar so +the newest request can start before the old solve deadline. diff --git a/go/internal/mpc/optimizer_transport.go b/go/internal/mpc/optimizer_transport.go index 233045cf0..82f6d3d34 100644 --- a/go/internal/mpc/optimizer_transport.go +++ b/go/internal/mpc/optimizer_transport.go @@ -307,41 +307,106 @@ func (g *contextGate) Unlock() { type UnixTransport struct{ socketPath string } +const unixCancelTimeout = 250 * time.Millisecond + +type unixCancelRequest struct { + Type string `json:"type"` + RequestID string `json:"request_id"` + ProtocolVersion int `json:"protocol_version"` +} + func NewUnixTransport(socketPath string) *UnixTransport { return &UnixTransport{socketPath: socketPath} } -func (t *UnixTransport) exchange(ctx context.Context, payload []byte) ([]byte, error) { +func (t *UnixTransport) exchange(ctx context.Context, payload []byte) ([]byte, bool, error) { var d net.Dialer conn, err := d.DialContext(ctx, "unix", t.socketPath) if err != nil { - return nil, fmt.Errorf("dial %s: %w", t.socketPath, err) + return nil, false, fmt.Errorf("dial %s: %w", t.socketPath, err) } defer conn.Close() if deadline, ok := ctx.Deadline(); ok { - _ = conn.SetDeadline(deadline) + // scanLine owns read cancellation. A read deadline could race ctx.Done + // and turn the caller's deadline into an unrelated I/O timeout. + _ = conn.SetWriteDeadline(deadline) } if _, err := conn.Write(append(append([]byte(nil), payload...), '\n')); err != nil { - return nil, fmt.Errorf("write unix optimizer: %w", err) + return nil, false, fmt.Errorf("write unix optimizer: %w", err) } scanner := bufio.NewScanner(conn) scanner.Buffer(make([]byte, 64*1024), 16*1024*1024) - return scanLine(ctx, scanner) + line, err := scanLine(ctx, scanner) + return line, true, err } func (t *UnixTransport) RoundTrip(ctx context.Context, payload []byte) ([]byte, error) { - return t.exchange(ctx, payload) + line, sent, err := t.exchange(ctx, payload) + if err == nil { + return line, nil + } + ctxErr := ctx.Err() + if !sent || ctxErr == nil { + return nil, err + } + requestID := optimizerRequestID(payload) + if requestID == "" { + return nil, err + } + + // The request may still be running in the shared sidecar after its caller + // has gone away. A fresh connection lets a current worker interrupt it; + // older workers reject the unknown frame without changing this error path. + _ = t.cancelRequest(requestID) + return nil, ctxErr } func (t *UnixTransport) Health(ctx context.Context) (OptimizerRuntimeInfo, error) { payload, _ := json.Marshal(map[string]any{"type": "handshake", "protocol_version": OptimizerProtocolVersion}) - line, err := t.exchange(ctx, payload) + line, _, err := t.exchange(ctx, payload) if err != nil { return OptimizerRuntimeInfo{}, err } return decodeOptimizerHandshake(line, "unix") } +func optimizerRequestID(payload []byte) string { + var request struct { + RequestID string `json:"request_id"` + } + if json.Unmarshal(payload, &request) != nil { + return "" + } + return request.RequestID +} + +func (t *UnixTransport) cancelRequest(requestID string) error { + ctx, cancel := context.WithTimeout(context.Background(), unixCancelTimeout) + defer cancel() + + var d net.Dialer + conn, err := d.DialContext(ctx, "unix", t.socketPath) + if err != nil { + return err + } + defer conn.Close() + if deadline, ok := ctx.Deadline(); ok { + _ = conn.SetWriteDeadline(deadline) + } + payload, err := json.Marshal(unixCancelRequest{ + Type: "cancel_request", + RequestID: requestID, + ProtocolVersion: OptimizerProtocolVersion, + }) + if err != nil { + return err + } + if _, err := conn.Write(append(payload, '\n')); err != nil { + return fmt.Errorf("write unix optimizer cancellation: %w", err) + } + return nil +} + func decodeOptimizerHandshake(line []byte, transport string) (OptimizerRuntimeInfo, error) { var info OptimizerRuntimeInfo if err := json.Unmarshal(line, &info); err != nil { diff --git a/go/internal/mpc/optimizer_transport_test.go b/go/internal/mpc/optimizer_transport_test.go index 7a7c17843..3db17a54d 100644 --- a/go/internal/mpc/optimizer_transport_test.go +++ b/go/internal/mpc/optimizer_transport_test.go @@ -131,6 +131,307 @@ func TestUnixTransportHandshakeAndRoundTrip(t *testing.T) { } } +func TestUnixTransportCancelsSentRequestOnCallerCancellation(t *testing.T) { + path := fmt.Sprintf("/tmp/ftw-opt-cancel-%d.sock", time.Now().UnixNano()) + t.Cleanup(func() { _ = os.Remove(path) }) + listener, err := net.Listen("unix", path) + if err != nil { + t.Fatal(err) + } + defer listener.Close() + + requestRead := make(chan []byte, 1) + cancelRead := make(chan []byte, 1) + serverDone := make(chan error, 1) + go func() { + requestConn, err := listener.Accept() + if err != nil { + serverDone <- err + return + } + defer requestConn.Close() + requestScanner := bufio.NewScanner(requestConn) + if !requestScanner.Scan() { + serverDone <- requestScanner.Err() + return + } + requestRead <- append([]byte(nil), requestScanner.Bytes()...) + + cancelConn, err := listener.Accept() + if err != nil { + serverDone <- err + return + } + defer cancelConn.Close() + cancelScanner := bufio.NewScanner(cancelConn) + if !cancelScanner.Scan() { + serverDone <- cancelScanner.Err() + return + } + cancelRead <- append([]byte(nil), cancelScanner.Bytes()...) + serverDone <- nil + }() + + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + payload := []byte(`{"schema_version":1,"request_id":"plan-42"}`) + go func() { + _, err := NewUnixTransport(path).RoundTrip(ctx, payload) + result <- err + }() + + select { + case got := <-requestRead: + if string(got) != string(payload) { + t.Fatalf("request = %s, want %s", got, payload) + } + case <-time.After(time.Second): + t.Fatal("sidecar did not receive the optimizer request") + } + cancel() + + select { + case err := <-result: + if !errors.Is(err, context.Canceled) { + t.Fatalf("RoundTrip error = %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("RoundTrip did not return after cancellation") + } + + var frame map[string]any + select { + case raw := <-cancelRead: + if err := json.Unmarshal(raw, &frame); err != nil { + t.Fatalf("decode cancel frame %q: %v", raw, err) + } + case <-time.After(time.Second): + t.Fatal("sidecar did not receive a cancel request") + } + if len(frame) != 3 || frame["type"] != "cancel_request" || frame["request_id"] != "plan-42" || frame["protocol_version"] != float64(OptimizerProtocolVersion) { + t.Fatalf("cancel frame = %#v", frame) + } + if err := <-serverDone; err != nil { + t.Fatalf("sidecar server: %v", err) + } +} + +func TestUnixTransportCancelsSentRequestOnCallerDeadline(t *testing.T) { + path := fmt.Sprintf("/tmp/ftw-opt-deadline-%d.sock", time.Now().UnixNano()) + t.Cleanup(func() { _ = os.Remove(path) }) + listener, err := net.Listen("unix", path) + if err != nil { + t.Fatal(err) + } + defer listener.Close() + + requestRead := make(chan struct{}) + cancelRead := make(chan []byte, 1) + serverDone := make(chan error, 1) + go func() { + requestConn, err := listener.Accept() + if err != nil { + serverDone <- err + return + } + defer requestConn.Close() + requestScanner := bufio.NewScanner(requestConn) + if !requestScanner.Scan() { + serverDone <- requestScanner.Err() + return + } + close(requestRead) + + cancelConn, err := listener.Accept() + if err != nil { + serverDone <- err + return + } + defer cancelConn.Close() + cancelScanner := bufio.NewScanner(cancelConn) + if !cancelScanner.Scan() { + serverDone <- cancelScanner.Err() + return + } + cancelRead <- append([]byte(nil), cancelScanner.Bytes()...) + serverDone <- nil + }() + + ctx, cancel := context.WithTimeout(context.Background(), 250*time.Millisecond) + defer cancel() + result := make(chan error, 1) + go func() { + _, err := NewUnixTransport(path).RoundTrip(ctx, []byte(`{"request_id":"deadline-plan"}`)) + result <- err + }() + select { + case <-requestRead: + case <-time.After(time.Second): + t.Fatal("sidecar did not receive the optimizer request") + } + select { + case err := <-result: + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("RoundTrip error = %v, want context.DeadlineExceeded", err) + } + case <-time.After(time.Second): + t.Fatal("RoundTrip did not return after its deadline") + } + + select { + case raw := <-cancelRead: + var frame unixCancelRequest + if err := json.Unmarshal(raw, &frame); err != nil { + t.Fatalf("decode cancel frame %q: %v", raw, err) + } + if frame.Type != "cancel_request" || frame.RequestID != "deadline-plan" || frame.ProtocolVersion != OptimizerProtocolVersion { + t.Fatalf("cancel frame = %+v", frame) + } + case <-time.After(time.Second): + t.Fatal("sidecar did not receive a cancel request") + } + if err := <-serverDone; err != nil { + t.Fatalf("sidecar server: %v", err) + } +} + +func TestUnixTransportDoesNotCancelRequestWithoutID(t *testing.T) { + path := fmt.Sprintf("/tmp/ftw-opt-no-id-%d.sock", time.Now().UnixNano()) + t.Cleanup(func() { _ = os.Remove(path) }) + listener, err := net.ListenUnix("unix", &net.UnixAddr{Name: path, Net: "unix"}) + if err != nil { + t.Fatal(err) + } + defer listener.Close() + + requestRead := make(chan struct{}) + checkSecond := make(chan struct{}) + serverDone := make(chan error, 1) + go func() { + requestConn, err := listener.Accept() + if err != nil { + serverDone <- err + return + } + defer requestConn.Close() + requestScanner := bufio.NewScanner(requestConn) + if !requestScanner.Scan() { + serverDone <- requestScanner.Err() + return + } + close(requestRead) + <-checkSecond + if err := listener.SetDeadline(time.Now().Add(100 * time.Millisecond)); err != nil { + serverDone <- err + return + } + second, err := listener.Accept() + if err == nil { + _ = second.Close() + serverDone <- errors.New("received unexpected cancel connection") + return + } + if timeout, ok := err.(net.Error); !ok || !timeout.Timeout() { + serverDone <- err + return + } + serverDone <- nil + }() + + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + _, err := NewUnixTransport(path).RoundTrip(ctx, []byte(`{"schema_version":1}`)) + result <- err + }() + select { + case <-requestRead: + case <-time.After(time.Second): + t.Fatal("sidecar did not receive the optimizer request") + } + cancel() + select { + case err := <-result: + if !errors.Is(err, context.Canceled) { + t.Fatalf("RoundTrip error = %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("RoundTrip did not return after cancellation") + } + close(checkSecond) + if err := <-serverDone; err != nil { + t.Fatal(err) + } +} + +func TestUnixTransportDoesNotCancelAfterOrdinaryReadFailure(t *testing.T) { + path := fmt.Sprintf("/tmp/ftw-opt-read-failure-%d.sock", time.Now().UnixNano()) + t.Cleanup(func() { _ = os.Remove(path) }) + listener, err := net.ListenUnix("unix", &net.UnixAddr{Name: path, Net: "unix"}) + if err != nil { + t.Fatal(err) + } + defer listener.Close() + + requestRead := make(chan struct{}) + checkSecond := make(chan struct{}) + serverDone := make(chan error, 1) + go func() { + requestConn, err := listener.Accept() + if err != nil { + serverDone <- err + return + } + requestScanner := bufio.NewScanner(requestConn) + if !requestScanner.Scan() { + _ = requestConn.Close() + serverDone <- requestScanner.Err() + return + } + close(requestRead) + _ = requestConn.Close() + <-checkSecond + if err := listener.SetDeadline(time.Now().Add(100 * time.Millisecond)); err != nil { + serverDone <- err + return + } + second, err := listener.Accept() + if err == nil { + _ = second.Close() + serverDone <- errors.New("received unexpected cancel connection") + return + } + if timeout, ok := err.(net.Error); !ok || !timeout.Timeout() { + serverDone <- err + return + } + serverDone <- nil + }() + + result := make(chan error, 1) + go func() { + _, err := NewUnixTransport(path).RoundTrip(context.Background(), []byte(`{"request_id":"plan-42"}`)) + result <- err + }() + select { + case <-requestRead: + case <-time.After(time.Second): + t.Fatal("sidecar did not receive the optimizer request") + } + select { + case err := <-result: + if err == nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("RoundTrip error = %v, want ordinary read failure", err) + } + case <-time.After(time.Second): + t.Fatal("RoundTrip did not return after the sidecar closed") + } + close(checkSecond) + if err := <-serverDone; err != nil { + t.Fatal(err) + } +} + func TestProcessTransportHealthPerformsCompatibleHandshake(t *testing.T) { if len(os.Args) > 0 && os.Args[len(os.Args)-1] == "process-health-helper" { scanner := bufio.NewScanner(os.Stdin) diff --git a/optimizer/ftw_optimizer/deadline.py b/optimizer/ftw_optimizer/deadline.py index 47a9aa9ca..afe428204 100644 --- a/optimizer/ftw_optimizer/deadline.py +++ b/optimizer/ftw_optimizer/deadline.py @@ -1,5 +1,6 @@ from __future__ import annotations +import threading import time from collections.abc import Callable from dataclasses import dataclass, field @@ -12,6 +13,17 @@ class SolveDeadlineExceeded(RuntimeError): """The request's one worker-side time budget has been spent.""" +class SolveCancelled(SolveDeadlineExceeded): + """The caller cancelled the request before it could publish a result.""" + + +@dataclass +class _CancellationState: + cancelled: threading.Event = field(default_factory=threading.Event) + lock: threading.Lock = field(default_factory=threading.Lock) + active_highs: Any | None = None + + @dataclass(frozen=True) class SolveDeadline: expires_at: float @@ -20,6 +32,11 @@ class SolveDeadline: repr=False, compare=False, ) + _cancellation: _CancellationState = field( + default_factory=_CancellationState, + repr=False, + compare=False, + ) @classmethod def from_payload( @@ -39,6 +56,8 @@ def from_payload( return cls(started_at + budget_s, clock) def remaining_s(self, phase: str = "optimizer request") -> float: + if self.is_cancelled(): + raise SolveCancelled(f"{phase} was cancelled") remaining = self.expires_at - self.clock() if remaining <= 0.0: raise SolveDeadlineExceeded(f"{phase} deadline exceeded") @@ -46,3 +65,24 @@ def remaining_s(self, phase: str = "optimizer request") -> float: def check(self, phase: str = "optimizer request") -> None: self.remaining_s(phase) + + def cancel(self) -> None: + self._cancellation.cancelled.set() + with self._cancellation.lock: + highs = self._cancellation.active_highs + if highs is not None: + highs.cancelSolve() + + def is_cancelled(self) -> bool: + return self._cancellation.cancelled.is_set() + + def attach_highs(self, highs: Any) -> None: + with self._cancellation.lock: + if self._cancellation.active_highs is not None: + raise RuntimeError("a HiGHS solve is already attached") + self._cancellation.active_highs = highs + + def detach_highs(self, highs: Any) -> None: + with self._cancellation.lock: + if self._cancellation.active_highs is highs: + self._cancellation.active_highs = None diff --git a/optimizer/ftw_optimizer/direct_highs.py b/optimizer/ftw_optimizer/direct_highs.py index dd71c4bc7..f7400cd07 100644 --- a/optimizer/ftw_optimizer/direct_highs.py +++ b/optimizer/ftw_optimizer/direct_highs.py @@ -9,7 +9,7 @@ import numpy as np from . import SCHEMA_VERSION -from .deadline import SolveDeadline, SolveDeadlineExceeded +from .deadline import SolveCancelled, SolveDeadline, SolveDeadlineExceeded from .model import ( _arbitrage_spread_ore_kwh, _solver_options, @@ -924,12 +924,42 @@ def _run_optimal( phase: str, deadline: SolveDeadline | float, ) -> None: - run_status = highs.run() + _remaining_time_s(deadline) + if isinstance(deadline, SolveDeadline): + if not highs.HandleUserInterrupt: + highs.HandleUserInterrupt = True + deadline.attach_highs(highs) + try: + try: + deadline.check(f"direct HiGHS {phase} solve") + solver_thread = highs.startSolve() + # startSolve resets HiGHS' stop flag. Repeat a cancellation that + # arrived after attachment but before the solver thread started. + if deadline.is_cancelled(): + highs.cancelSolve() + run_status = highs.joinSolve(solver_thread) + except Exception: + # Cancellation or expiry is the request result even when HiGHS + # reports its own concurrent start/join error first. + deadline.check(f"direct HiGHS {phase} solve") + raise + finally: + deadline.detach_highs(highs) + deadline.check(f"direct HiGHS {phase} solve") + else: + run_status = highs.run() status = highs.getModelStatus() + if status in { + highspy.HighsModelStatus.kInterrupt, + highspy.HighsModelStatus.kHighsInterrupt, + }: + raise SolveCancelled(f"direct HiGHS {phase} solve was cancelled") if status == highspy.HighsModelStatus.kTimeLimit: raise SolveDeadlineExceeded( f"direct HiGHS {phase} solve deadline exceeded" ) + if run_status is None: + raise DirectHighsError(f"HiGHS {phase} solve returned no status") _require_ok(run_status, f"run {phase} solve") if status != highspy.HighsModelStatus.kOptimal: raise DirectHighsError(f"HiGHS {phase} solve failed with status {status}") diff --git a/optimizer/ftw_optimizer/worker.py b/optimizer/ftw_optimizer/worker.py index 725dc0b62..be4f70daa 100644 --- a/optimizer/ftw_optimizer/worker.py +++ b/optimizer/ftw_optimizer/worker.py @@ -11,13 +11,14 @@ import threading import time import traceback +from collections import OrderedDict from collections.abc import Callable from pathlib import Path from typing import Any import cvxpy as cp -from .deadline import SolveDeadline, SolveDeadlineExceeded +from .deadline import SolveCancelled, SolveDeadline, SolveDeadlineExceeded from .model import solve from .protocol import ParsedRequest, ProtocolError, error_response, parse_request @@ -32,8 +33,95 @@ # the window stop using this optimizer at once. MIN_PROTOCOL_VERSION = 1 PROTOCOL_VERSION = 1 -FEATURES = ["champion", "recourse", "multistage", "commercial_constraints_v1"] -SOLVE_LOCK = threading.Lock() +FEATURES = [ + "champion", + "recourse", + "multistage", + "commercial_constraints_v1", + "cancel_request", +] + + +class _SolveLock: + def __init__(self) -> None: + self._condition = threading.Condition() + self._held = False + + def acquire_until(self, deadline: SolveDeadline) -> bool: + with self._condition: + while self._held: + self._condition.wait( + timeout=min( + deadline.remaining_s("optimizer queue"), + threading.TIMEOUT_MAX, + ) + ) + deadline.check("optimizer queue") + self._held = True + return True + + def release(self) -> None: + with self._condition: + if not self._held: + raise RuntimeError("cannot release an unlocked solve lock") + self._held = False + self._condition.notify_all() + + def locked(self) -> bool: + with self._condition: + return self._held + + def notify_waiters(self) -> None: + with self._condition: + self._condition.notify_all() + + +class _ActiveRequests: + def __init__(self, max_pending_cancels: int = 256) -> None: + self._lock = threading.Lock() + self._active: dict[str, list[SolveDeadline]] = {} + self._pending_cancels: OrderedDict[str, None] = OrderedDict() + self._max_pending_cancels = max_pending_cancels + + def register(self, request_id: str, deadline: SolveDeadline) -> None: + with self._lock: + self._active.setdefault(request_id, []).append(deadline) + cancel_now = request_id in self._pending_cancels + self._pending_cancels.pop(request_id, None) + if cancel_now: + deadline.cancel() + + def unregister(self, request_id: str, deadline: SolveDeadline) -> None: + with self._lock: + deadlines = self._active.get(request_id) + if deadlines is None: + return + self._active[request_id] = [ + candidate for candidate in deadlines if candidate is not deadline + ] + if not self._active[request_id]: + del self._active[request_id] + + def cancel(self, request_id: str) -> bool: + with self._lock: + deadlines = tuple(self._active.get(request_id, ())) + if not deadlines: + self._pending_cancels[request_id] = None + self._pending_cancels.move_to_end(request_id) + while len(self._pending_cancels) > self._max_pending_cancels: + self._pending_cancels.popitem(last=False) + for deadline in deadlines: + try: + deadline.cancel() + except Exception: + # The token was set before HiGHS was asked to stop. Keep the + # cancel connection alive even if that best-effort call fails. + traceback.print_exc(file=sys.stderr) + return bool(deadlines) + + +SOLVE_LOCK: Any = _SolveLock() +ACTIVE_REQUESTS = _ActiveRequests() def release_unused_memory() -> None: @@ -75,6 +163,8 @@ def handle( return response except ProtocolError as exc: return error_response(request_id, "invalid_request", str(exc)) + except SolveCancelled as exc: + return error_response(request_id, "cancelled", str(exc)) except SolveDeadlineExceeded as exc: return error_response(request_id, "deadline_exceeded", str(exc)) except cp.error.SolverError as exc: @@ -102,6 +192,52 @@ def handshake(raw: Any) -> dict[str, Any] | None: } +def cancel_request(raw: Any) -> dict[str, Any] | None: + if not isinstance(raw, dict) or raw.get("type") != "cancel_request": + return None + request_id = raw.get("request_id") + if not isinstance(request_id, str) or not request_id: + return error_response( + "unknown", + "invalid_request", + "request_id must be a non-empty string", + ) + protocol_version = raw.get("protocol_version", PROTOCOL_VERSION) + if ( + isinstance(protocol_version, bool) + or not isinstance(protocol_version, int) + or not MIN_PROTOCOL_VERSION <= protocol_version <= PROTOCOL_VERSION + ): + return error_response( + request_id, + "invalid_request", + f"unsupported protocol_version {protocol_version!r}; expected " + f"{MIN_PROTOCOL_VERSION}..{PROTOCOL_VERSION}", + ) + active = ACTIVE_REQUESTS.cancel(request_id) + notify_waiters = getattr(SOLVE_LOCK, "notify_waiters", None) + if notify_waiters is not None: + notify_waiters() + return { + "type": "cancel_ack", + "protocol_version": PROTOCOL_VERSION, + "request_id": request_id, + "ok": True, + "active": active, + } + + +def _acquire_solve_lock(deadline: SolveDeadline) -> bool: + acquire_until = getattr(SOLVE_LOCK, "acquire_until", None) + if acquire_until is not None: + return bool(acquire_until(deadline)) + wait_s = min( + deadline.remaining_s("optimizer queue"), + threading.TIMEOUT_MAX, + ) + return bool(SOLVE_LOCK.acquire(timeout=wait_s)) + + def process_stream( reader: Any, writer: Any, @@ -119,69 +255,82 @@ def process_stream( else: response = handshake(raw) if response is None: - # Handshakes stay responsive while a solve is in progress, - # but solver state and its memory cleanup remain serialized. + response = cancel_request(raw) + if response is None: + # Handshakes stay responsive while a solve is in progress. + # Cancel frames also bypass the solve lock so they can stop its + # current owner or remove a queued request at once. request_id = "unknown" + deadline: SolveDeadline | None = None + registered = False try: - parsed = parse_request(raw) - request_id = parsed.request_id - deadline = SolveDeadline.from_payload( - parsed.payload, - started_at=received_at, - clock=clock, - ) - wait_s = min( - deadline.remaining_s("optimizer queue"), - threading.TIMEOUT_MAX, - ) - except ProtocolError as exc: - response = error_response(request_id, "invalid_request", str(exc)) - except SolveDeadlineExceeded as exc: - response = error_response( - request_id, - "deadline_exceeded", - str(exc), - ) - else: - if not SOLVE_LOCK.acquire(timeout=wait_s): + try: + parsed = parse_request(raw) + request_id = parsed.request_id + deadline = SolveDeadline.from_payload( + parsed.payload, + started_at=received_at, + clock=clock, + ) + ACTIVE_REQUESTS.register(request_id, deadline) + registered = True + acquired = _acquire_solve_lock(deadline) + except ProtocolError as exc: response = error_response( - parsed.request_id, - "deadline_exceeded", - "optimizer queue deadline exceeded", + request_id, + "invalid_request", + str(exc), ) - writer.write( - json.dumps( - response, - separators=(",", ":"), - allow_nan=False, - ) - + "\n" + except SolveCancelled as exc: + response = error_response( + request_id, + "cancelled", + str(exc), ) - writer.flush() - continue - try: - response = handle( - raw, - received_at=received_at, - clock=clock, - parsed=parsed, - deadline=deadline, + except SolveDeadlineExceeded as exc: + response = error_response( + request_id, + "deadline_exceeded", + str(exc), ) - try: - writer.write( - json.dumps( - response, - separators=(",", ":"), - allow_nan=False, - ) - + "\n" + else: + if not acquired: + response = error_response( + request_id, + "deadline_exceeded", + "optimizer queue deadline exceeded", ) - writer.flush() - finally: - response = None - release_unused_memory() - finally: - SOLVE_LOCK.release() + else: + try: + response = handle( + raw, + received_at=received_at, + clock=clock, + parsed=parsed, + deadline=deadline, + ) + try: + if not deadline.is_cancelled(): + writer.write( + json.dumps( + response, + separators=(",", ":"), + allow_nan=False, + ) + + "\n" + ) + writer.flush() + finally: + response = None + release_unused_memory() + finally: + SOLVE_LOCK.release() + continue + finally: + if registered: + assert deadline is not None + ACTIVE_REQUESTS.unregister(request_id, deadline) + if deadline is not None and deadline.is_cancelled(): continue writer.write(json.dumps(response, separators=(",", ":"), allow_nan=False) + "\n") writer.flush() diff --git a/optimizer/tests/test_deadline.py b/optimizer/tests/test_deadline.py index c8f95fb7f..66fb9f419 100644 --- a/optimizer/tests/test_deadline.py +++ b/optimizer/tests/test_deadline.py @@ -1,11 +1,17 @@ from __future__ import annotations +import threading + import cvxpy as cp import highspy import pytest from ftw_optimizer import shared_highs -from ftw_optimizer.deadline import SolveDeadline, SolveDeadlineExceeded +from ftw_optimizer.deadline import ( + SolveCancelled, + SolveDeadline, + SolveDeadlineExceeded, +) from ftw_optimizer.direct_highs import ( DirectHighsError, _remaining_time_s, @@ -88,10 +94,17 @@ def __init__( ) -> None: self.status = status self.run_status = run_status + self.HandleUserInterrupt = False def run(self) -> highspy.HighsStatus: return self.run_status + def startSolve(self) -> object: + return object() + + def joinSolve(self, _solver_thread: object) -> highspy.HighsStatus: + return self.run_status + def getModelStatus(self) -> highspy.HighsModelStatus: return self.status @@ -115,3 +128,216 @@ def test_direct_highs_time_limit_is_a_deadline_not_a_fallback_error() -> None: "service", deadline, ) + + +class BlockingHighs: + def __init__(self) -> None: + self.HandleUserInterrupt = False + self.start_entered = threading.Event() + self.allow_start = threading.Event() + self.cancelled = threading.Event() + self.cancel_calls = 0 + + def startSolve(self) -> object: + self.start_entered.set() + if not self.allow_start.wait(timeout=1): + raise TimeoutError("test did not allow HiGHS to start") + # HiGHS clears its stop flag in startSolve. + self.cancelled.clear() + return object() + + def cancelSolve(self) -> None: + self.cancel_calls += 1 + self.cancelled.set() + + def joinSolve(self, _solver_thread: object) -> highspy.HighsStatus: + if not self.cancelled.wait(timeout=1): + raise TimeoutError("test cancellation did not reach HiGHS") + return highspy.HighsStatus.kWarning + + def getModelStatus(self) -> highspy.HighsModelStatus: + return highspy.HighsModelStatus.kInterrupt + + +class CountingInterruptHighs(FakeHighs): + def __init__(self) -> None: + self._handle_user_interrupt = False + self.interrupt_enable_calls = 0 + super().__init__(highspy.HighsModelStatus.kOptimal) + + @property + def HandleUserInterrupt(self) -> bool: + return self._handle_user_interrupt + + @HandleUserInterrupt.setter + def HandleUserInterrupt(self, enabled: bool) -> None: + self._handle_user_interrupt = enabled + if enabled: + self.interrupt_enable_calls += 1 + + +class RunningHighs: + def __init__(self) -> None: + self.HandleUserInterrupt = False + self.join_entered = threading.Event() + self.cancelled = threading.Event() + self.cancel_calls = 0 + + def startSolve(self) -> object: + return object() + + def cancelSolve(self) -> None: + self.cancel_calls += 1 + self.cancelled.set() + + def joinSolve(self, _solver_thread: object) -> highspy.HighsStatus: + self.join_entered.set() + if not self.cancelled.wait(timeout=1): + raise TimeoutError("test cancellation did not reach HiGHS") + return highspy.HighsStatus.kWarning + + def getModelStatus(self) -> highspy.HighsModelStatus: + return highspy.HighsModelStatus.kHighsInterrupt + + +class RaisingOnCancelHighs(RunningHighs): + def joinSolve(self, _solver_thread: object) -> highspy.HighsStatus: + self.join_entered.set() + if not self.cancelled.wait(timeout=1): + raise TimeoutError("test cancellation did not reach HiGHS") + raise RuntimeError("HiGHS join failed during cancellation") + + +class RaisingAfterDeadlineHighs(FakeHighs): + def __init__(self, clock: FakeClock) -> None: + super().__init__(highspy.HighsModelStatus.kSolveError) + self.clock = clock + + def joinSolve(self, _solver_thread: object) -> highspy.HighsStatus: + self.clock.advance(2.0) + raise RuntimeError("HiGHS join failed after the deadline") + + +def test_direct_highs_enables_interrupt_callbacks_once_per_model() -> None: + deadline = SolveDeadline(1.0, FakeClock()) + highs = CountingInterruptHighs() + + _run_optimal(highs, "service", deadline) + _run_optimal(highs, "economic", deadline) + + assert highs.interrupt_enable_calls == 1 + + +def test_cancel_interrupts_an_active_direct_highs_solve() -> None: + deadline = SolveDeadline(1.0, FakeClock()) + highs = RunningHighs() + errors: list[BaseException] = [] + + def run() -> None: + try: + _run_optimal(highs, "service", deadline) + except BaseException as exc: + errors.append(exc) + + thread = threading.Thread(target=run) + thread.start() + assert highs.join_entered.wait(timeout=1) + + deadline.cancel() + thread.join(timeout=1) + + assert not thread.is_alive() + assert len(errors) == 1 + assert isinstance(errors[0], SolveCancelled) + assert highs.cancel_calls == 1 + + +def test_cancellation_wins_when_highs_join_raises() -> None: + deadline = SolveDeadline(1.0, FakeClock()) + highs = RaisingOnCancelHighs() + errors: list[BaseException] = [] + + def run() -> None: + try: + _run_optimal(highs, "service", deadline) + except BaseException as exc: + errors.append(exc) + + thread = threading.Thread(target=run) + thread.start() + assert highs.join_entered.wait(timeout=1) + + deadline.cancel() + thread.join(timeout=1) + + assert not thread.is_alive() + assert len(errors) == 1 + assert isinstance(errors[0], SolveCancelled) + assert isinstance(errors[0].__context__, RuntimeError) + + +def test_deadline_wins_when_highs_join_raises_after_expiry() -> None: + clock = FakeClock() + deadline = SolveDeadline(1.0, clock) + + with pytest.raises(SolveDeadlineExceeded, match="deadline exceeded"): + _run_optimal(RaisingAfterDeadlineHighs(clock), "service", deadline) + + +def test_direct_highs_repeats_cancel_after_start_resets_the_stop_flag() -> None: + deadline = SolveDeadline(1.0, FakeClock()) + highs = BlockingHighs() + errors: list[BaseException] = [] + phase_two_started = threading.Event() + + def run() -> None: + try: + _run_optimal(highs, "service", deadline) + phase_two_started.set() + except BaseException as exc: + errors.append(exc) + + thread = threading.Thread(target=run) + thread.start() + assert highs.start_entered.wait(timeout=1) + + deadline.cancel() + highs.allow_start.set() + thread.join(timeout=1) + + assert not thread.is_alive() + assert len(errors) == 1 + assert isinstance(errors[0], SolveCancelled) + assert highs.cancel_calls == 2 + assert not phase_two_started.is_set() + + +def test_cancelled_direct_highs_does_not_fall_back_to_cvxpy(monkeypatch) -> None: + deadline = SolveDeadline(1.0, FakeClock()) + direct_calls: list[SolveDeadline] = [] + + def cancel_direct( + _payload: dict, + _started: float, + received_deadline: SolveDeadline, + ) -> dict: + direct_calls.append(received_deadline) + raise SolveCancelled("direct HiGHS solve was cancelled") + + monkeypatch.setattr(shared_highs, "solve_shared_highs", cancel_direct) + + with pytest.raises(SolveCancelled, match="cancelled"): + solve( + { + "settings": { + "shared_backend": "auto", + "time_limit_s": 10.0, + }, + "commercial_constraints": {}, + "slots": [{}], + "storages": [], + }, + deadline=deadline, + ) + + assert direct_calls == [deadline] diff --git a/optimizer/tests/test_worker.py b/optimizer/tests/test_worker.py index abbcdbde9..ec9f9e713 100644 --- a/optimizer/tests/test_worker.py +++ b/optimizer/tests/test_worker.py @@ -4,7 +4,15 @@ import json import threading +import pytest + from ftw_optimizer import worker +from ftw_optimizer.deadline import SolveDeadline + + +@pytest.fixture(autouse=True) +def reset_active_requests(monkeypatch) -> None: + monkeypatch.setattr(worker, "ACTIVE_REQUESTS", worker._ActiveRequests()) def test_health_stays_responsive_without_cleaning_memory_during_solve( @@ -56,7 +64,9 @@ def run_solve() -> None: health_thread.start() health_thread.join(timeout=1) assert not health_thread.is_alive() - assert '"name":"ftw-optimizer"' in health_output.getvalue() + health = json.loads(health_output.getvalue()) + assert health["name"] == "ftw-optimizer" + assert "cancel_request" in health["features"] assert cleanup_calls == [] finish_solve.set() @@ -103,6 +113,16 @@ def acquire(self, blocking: bool = True, timeout: float = -1) -> bool: self.held = True return True + def acquire_until(self, deadline: SolveDeadline) -> bool: + with self.clock.condition: + while self.held: + self.queued.set() + deadline.check("optimizer queue") + self.clock.condition.wait() + deadline.check("optimizer queue") + self.held = True + return True + def release(self) -> None: with self.clock.condition: self.held = False @@ -112,6 +132,10 @@ def locked(self) -> bool: with self.clock.condition: return self.held + def notify_waiters(self) -> None: + with self.clock.condition: + self.clock.condition.notify_all() + def test_expired_request_leaves_solve_queue_without_running( monkeypatch, @@ -214,3 +238,303 @@ def fake_solve(_payload: dict, **_kwargs: object) -> dict[str, object]: assert solve_calls == 1 assert response["error"]["code"] == "deadline_exceeded" + + +def request_stream(request_id: str, budget_s: float = 10.0) -> io.StringIO: + return io.StringIO( + json.dumps( + { + "schema_version": 1, + "request_id": request_id, + "settings": {"time_limit_s": budget_s}, + "slots": [{}], + } + ) + + "\n" + ) + + +def cancel_stream(request_id: str) -> io.StringIO: + return io.StringIO( + json.dumps( + { + "type": "cancel_request", + "protocol_version": 1, + "request_id": request_id, + } + ) + + "\n" + ) + + +def test_active_cancel_releases_the_next_request_before_the_old_deadline( + monkeypatch, +) -> None: + clock = FakeClock() + solve_lock = FakeSolveLock(clock) + first_started = threading.Event() + let_cancelled_solve_check_token = threading.Event() + second_started = threading.Event() + solve_calls: list[str] = [] + thread_errors: list[BaseException] = [] + monkeypatch.setattr(worker, "SOLVE_LOCK", solve_lock) + monkeypatch.setattr(worker, "release_unused_memory", lambda: None) + + def fake_solve( + payload: dict, + *, + deadline: SolveDeadline, + ) -> dict[str, object]: + request_id = str(payload["request_id"]) + solve_calls.append(request_id) + if request_id == "first": + first_started.set() + if not let_cancelled_solve_check_token.wait(timeout=1): + raise TimeoutError("test did not finish cancellation") + deadline.check("fake active solve") + else: + assert clock() == 0.0 + second_started.set() + return {"ok": True, "request_id": request_id} + + monkeypatch.setattr(worker, "solve", fake_solve) + + def run(request_id: str, output: io.StringIO) -> None: + try: + worker.process_stream(request_stream(request_id), output, clock=clock) + except BaseException as exc: + thread_errors.append(exc) + + first_output = io.StringIO() + first = threading.Thread(target=run, args=("first", first_output)) + first.start() + assert first_started.wait(timeout=1) + + second_output = io.StringIO() + second = threading.Thread(target=run, args=("second", second_output)) + second.start() + assert solve_lock.queued.wait(timeout=1) + + cancel_output = io.StringIO() + worker.process_stream(cancel_stream("first"), cancel_output, clock=clock) + let_cancelled_solve_check_token.set() + + first.join(timeout=1) + second.join(timeout=1) + assert not first.is_alive() + assert not second.is_alive() + assert second_started.is_set() + assert first_output.getvalue() == "" + assert json.loads(second_output.getvalue())["request_id"] == "second" + assert json.loads(cancel_output.getvalue())["active"] is True + assert solve_calls == ["first", "second"] + assert clock() == 0.0 + assert thread_errors == [] + + +def test_queued_cancel_leaves_without_waiting_for_the_active_request( + monkeypatch, +) -> None: + clock = FakeClock() + solve_lock = FakeSolveLock(clock) + active_started = threading.Event() + finish_active = threading.Event() + solve_calls: list[str] = [] + thread_errors: list[BaseException] = [] + monkeypatch.setattr(worker, "SOLVE_LOCK", solve_lock) + monkeypatch.setattr(worker, "release_unused_memory", lambda: None) + + def fake_solve(payload: dict, **_kwargs: object) -> dict[str, object]: + request_id = str(payload["request_id"]) + solve_calls.append(request_id) + if request_id == "active": + active_started.set() + if not finish_active.wait(timeout=1): + raise TimeoutError("test did not release active request") + return {"ok": True, "request_id": request_id} + + monkeypatch.setattr(worker, "solve", fake_solve) + + def run(request_id: str, output: io.StringIO) -> None: + try: + worker.process_stream(request_stream(request_id), output, clock=clock) + except BaseException as exc: + thread_errors.append(exc) + + active_output = io.StringIO() + active = threading.Thread(target=run, args=("active", active_output)) + active.start() + assert active_started.wait(timeout=1) + + queued_output = io.StringIO() + queued = threading.Thread(target=run, args=("queued", queued_output)) + queued.start() + assert solve_lock.queued.wait(timeout=1) + + cancel_output = io.StringIO() + worker.process_stream(cancel_stream("queued"), cancel_output, clock=clock) + queued.join(timeout=1) + + assert not queued.is_alive() + assert active.is_alive() + assert queued_output.getvalue() == "" + assert json.loads(cancel_output.getvalue())["active"] is True + assert solve_calls == ["active"] + + finish_active.set() + active.join(timeout=1) + assert not active.is_alive() + assert json.loads(active_output.getvalue())["request_id"] == "active" + assert thread_errors == [] + + +def test_early_cancel_prevents_the_request_from_entering_the_solver( + monkeypatch, +) -> None: + solve_calls: list[str] = [] + monkeypatch.setattr(worker, "SOLVE_LOCK", _TestSolveLock()) + monkeypatch.setattr(worker, "release_unused_memory", lambda: None) + monkeypatch.setattr( + worker, + "solve", + lambda payload, **_kwargs: solve_calls.append(str(payload["request_id"])), + ) + + cancel_output = io.StringIO() + worker.process_stream(cancel_stream("early"), cancel_output) + request_output = io.StringIO() + worker.process_stream(request_stream("early"), request_output) + + assert json.loads(cancel_output.getvalue())["active"] is False + assert request_output.getvalue() == "" + assert solve_calls == [] + + +def test_cancel_accepts_an_older_protocol_version_in_the_worker_window( + monkeypatch, +) -> None: + monkeypatch.setattr(worker, "PROTOCOL_VERSION", 2) + output = io.StringIO() + + worker.process_stream(cancel_stream("future-worker"), output) + + response = json.loads(output.getvalue()) + assert response["ok"] is True + assert response["protocol_version"] == 2 + + +@pytest.mark.parametrize("protocol_version", [True, 0, 2, "1"]) +def test_cancel_rejects_a_protocol_version_outside_the_worker_window( + protocol_version: object, +) -> None: + output = io.StringIO() + request = cancel_stream("invalid-version") + raw = json.loads(request.getvalue()) + raw["protocol_version"] = protocol_version + + worker.process_stream(io.StringIO(json.dumps(raw) + "\n"), output) + + assert json.loads(output.getvalue())["error"]["code"] == "invalid_request" + + +def test_wrong_request_id_does_not_cancel_the_active_request(monkeypatch) -> None: + active_started = threading.Event() + finish_active = threading.Event() + active_deadline: list[SolveDeadline] = [] + thread_errors: list[BaseException] = [] + monkeypatch.setattr(worker, "SOLVE_LOCK", _TestSolveLock()) + monkeypatch.setattr(worker, "release_unused_memory", lambda: None) + + def fake_solve( + payload: dict, + *, + deadline: SolveDeadline, + ) -> dict[str, object]: + active_deadline.append(deadline) + active_started.set() + if not finish_active.wait(timeout=1): + raise TimeoutError("test did not release active request") + deadline.check("fake active solve") + return {"ok": True, "request_id": str(payload["request_id"])} + + monkeypatch.setattr(worker, "solve", fake_solve) + output = io.StringIO() + + def run() -> None: + try: + worker.process_stream(request_stream("active"), output) + except BaseException as exc: + thread_errors.append(exc) + + thread = threading.Thread(target=run) + thread.start() + assert active_started.wait(timeout=1) + + cancel_output = io.StringIO() + worker.process_stream(cancel_stream("different"), cancel_output) + assert json.loads(cancel_output.getvalue())["active"] is False + assert not active_deadline[0].is_cancelled() + + finish_active.set() + thread.join(timeout=1) + assert not thread.is_alive() + assert json.loads(output.getvalue())["ok"] is True + assert thread_errors == [] + + +def test_pending_cancel_registry_has_a_fixed_bound() -> None: + registry = worker._ActiveRequests(max_pending_cancels=2) + registry.cancel("evicted") + registry.cancel("kept-one") + registry.cancel("kept-two") + evicted = SolveDeadline(1.0, FakeClock()) + kept_one = SolveDeadline(1.0, FakeClock()) + kept_two = SolveDeadline(1.0, FakeClock()) + + registry.register("evicted", evicted) + registry.register("kept-one", kept_one) + registry.register("kept-two", kept_two) + + assert not evicted.is_cancelled() + assert kept_one.is_cancelled() + assert kept_two.is_cancelled() + + +class _TestSolveLock: + def __init__(self) -> None: + self.lock = threading.Lock() + + def acquire_until(self, deadline: SolveDeadline) -> bool: + deadline.check("optimizer queue") + return self.lock.acquire() + + def release(self) -> None: + self.lock.release() + + def locked(self) -> bool: + return self.lock.locked() + + def notify_waiters(self) -> None: + return None + + +class _FailingCancelHighs: + def cancelSolve(self) -> None: + raise RuntimeError("cancel failed") + + +def test_cancel_frame_survives_a_highs_cancel_failure(monkeypatch, capsys) -> None: + monkeypatch.setattr(worker, "SOLVE_LOCK", _TestSolveLock()) + deadline = SolveDeadline(1.0, FakeClock()) + highs = _FailingCancelHighs() + deadline.attach_highs(highs) + worker.ACTIVE_REQUESTS.register("failing", deadline) + output = io.StringIO() + + worker.process_stream(cancel_stream("failing"), output) + + assert json.loads(output.getvalue())["active"] is True + assert deadline.is_cancelled() + assert "cancel failed" in capsys.readouterr().err + deadline.detach_highs(highs) + worker.ACTIVE_REQUESTS.unregister("failing", deadline)