diff --git a/.changeset/optimizer-deadline-budget.md b/.changeset/optimizer-deadline-budget.md new file mode 100644 index 00000000..a8a18475 --- /dev/null +++ b/.changeset/optimizer-deadline-budget.md @@ -0,0 +1,5 @@ +--- +"ftw": patch +--- + +Bound each optimizer request to one worker deadline so expired queued work no longer blocks newer plans or resets its budget between solver phases. diff --git a/go/internal/mpc/optimizer_transport.go b/go/internal/mpc/optimizer_transport.go index 7ee9bad3..233045cf 100644 --- a/go/internal/mpc/optimizer_transport.go +++ b/go/internal/mpc/optimizer_transport.go @@ -12,7 +12,6 @@ import ( "os/exec" "strconv" "strings" - "sync" "time" "github.com/srcfl/ftw/go/internal/optimizercontract" @@ -86,7 +85,7 @@ type ProcessTransportConfig struct { type ProcessTransport struct { cfg ProcessTransportConfig - mu sync.Mutex + mu *contextGate cmd *exec.Cmd stdin io.WriteCloser scanner *bufio.Scanner @@ -98,16 +97,22 @@ func NewProcessTransport(cfg ProcessTransportConfig) (*ProcessTransport, error) if len(cfg.Command) == 0 || strings.TrimSpace(cfg.Command[0]) == "" { return nil, errors.New("optimizer command is empty") } - return &ProcessTransport{cfg: cfg}, nil + return &ProcessTransport{cfg: cfg, mu: newContextGate()}, nil } func (t *ProcessTransport) RoundTrip(ctx context.Context, payload []byte) ([]byte, error) { - t.mu.Lock() - defer t.mu.Unlock() + if err := t.mu.acquire(ctx); err != nil { + return nil, err + } + defer t.mu.release() t.cancelIdleStopLocked() if err := t.ensureStartedLocked(); err != nil { return nil, err } + if err := ctx.Err(); err != nil { + t.scheduleIdleStopLocked() + return nil, err + } if _, err := t.stdin.Write(append(append([]byte(nil), payload...), '\n')); err != nil { t.stopLocked() return nil, fmt.Errorf("write optimizer request: %w", err) @@ -122,15 +127,18 @@ func (t *ProcessTransport) RoundTrip(ctx context.Context, payload []byte) ([]byt } func (t *ProcessTransport) Health(ctx context.Context) (OptimizerRuntimeInfo, error) { - t.mu.Lock() - defer t.mu.Unlock() - t.cancelIdleStopLocked() - if err := ctx.Err(); err != nil { + if err := t.mu.acquire(ctx); err != nil { return OptimizerRuntimeInfo{}, err } + defer t.mu.release() + t.cancelIdleStopLocked() if err := t.ensureStartedLocked(); err != nil { return OptimizerRuntimeInfo{}, err } + if err := ctx.Err(); err != nil { + t.scheduleIdleStopLocked() + return OptimizerRuntimeInfo{}, err + } payload, _ := json.Marshal(map[string]any{ "type": "handshake", "protocol_version": OptimizerProtocolVersion, @@ -216,8 +224,8 @@ func (t *ProcessTransport) scheduleIdleStopLocked() { t.cancelIdleStopLocked() var timer *time.Timer timer = time.AfterFunc(t.cfg.IdleTimeout, func() { - t.mu.Lock() - defer t.mu.Unlock() + _ = t.mu.acquire(context.Background()) + defer t.mu.release() if t.idleTimer != timer { return } @@ -239,8 +247,8 @@ func (t *ProcessTransport) stopLocked() { } func (t *ProcessTransport) Close() error { - t.mu.Lock() - defer t.mu.Unlock() + _ = t.mu.acquire(context.Background()) + defer t.mu.release() t.cancelIdleStopLocked() if t.cmd == nil { return nil @@ -256,6 +264,47 @@ func (t *ProcessTransport) Close() error { } } +// contextGate serializes access to the warm process without hiding queue wait +// from the caller's deadline. A request canceled while it waits never reaches +// stdin, so it cannot occupy the worker after its result has become useless. +type contextGate struct { + token chan struct{} +} + +func newContextGate() *contextGate { + g := &contextGate{token: make(chan struct{}, 1)} + g.token <- struct{}{} + return g +} + +func (g *contextGate) acquire(ctx context.Context) error { + if err := ctx.Err(); err != nil { + return err + } + select { + case <-ctx.Done(): + return ctx.Err() + case <-g.token: + if err := ctx.Err(); err != nil { + g.release() + return err + } + return nil + } +} + +func (g *contextGate) release() { + g.token <- struct{}{} +} + +func (g *contextGate) Lock() { + _ = g.acquire(context.Background()) +} + +func (g *contextGate) Unlock() { + g.release() +} + type UnixTransport struct{ socketPath string } func NewUnixTransport(socketPath string) *UnixTransport { diff --git a/go/internal/mpc/optimizer_transport_test.go b/go/internal/mpc/optimizer_transport_test.go index cc3746f1..7a7c1784 100644 --- a/go/internal/mpc/optimizer_transport_test.go +++ b/go/internal/mpc/optimizer_transport_test.go @@ -37,6 +37,56 @@ func TestOptimizerProtocolVersionKeepsContractAlias(t *testing.T) { } } +func TestContextGateDropsCanceledWaiter(t *testing.T) { + gate := newContextGate() + if err := gate.acquire(context.Background()); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + waiting := make(chan struct{}) + errCh := make(chan error, 1) + go func() { + close(waiting) + errCh <- gate.acquire(ctx) + }() + <-waiting + cancel() + + select { + case err := <-errCh: + if !errors.Is(err, context.Canceled) { + t.Fatalf("acquire error = %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("canceled waiter remained blocked") + } + + gate.release() + acquireCtx, acquireCancel := context.WithTimeout(context.Background(), time.Second) + defer acquireCancel() + if err := gate.acquire(acquireCtx); err != nil { + t.Fatalf("gate stayed occupied after canceled waiter: %v", err) + } + gate.release() +} + +func TestProcessTransportRejectsCanceledContextBeforeWorkerLookup(t *testing.T) { + transport, err := NewProcessTransport(ProcessTransportConfig{ + Command: []string{"ftw-worker-that-does-not-exist"}, + }) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err = transport.RoundTrip(ctx, []byte(`{}`)) + if !errors.Is(err, context.Canceled) { + t.Fatalf("RoundTrip error = %v, want context.Canceled", err) + } +} + func TestUnixTransportHandshakeAndRoundTrip(t *testing.T) { path := fmt.Sprintf("/tmp/ftw-opt-%d.sock", time.Now().UnixNano()) t.Cleanup(func() { _ = os.Remove(path) }) diff --git a/optimizer/ftw_optimizer/deadline.py b/optimizer/ftw_optimizer/deadline.py new file mode 100644 index 00000000..47a9aa9c --- /dev/null +++ b/optimizer/ftw_optimizer/deadline.py @@ -0,0 +1,48 @@ +from __future__ import annotations + +import time +from collections.abc import Callable +from dataclasses import dataclass, field +from typing import Any + +from .protocol import positive_number, require_dict + + +class SolveDeadlineExceeded(RuntimeError): + """The request's one worker-side time budget has been spent.""" + + +@dataclass(frozen=True) +class SolveDeadline: + expires_at: float + clock: Callable[[], float] = field( + default=time.perf_counter, + repr=False, + compare=False, + ) + + @classmethod + def from_payload( + cls, + payload: dict[str, Any], + *, + started_at: float | None = None, + clock: Callable[[], float] = time.perf_counter, + ) -> SolveDeadline: + settings = require_dict(payload.get("settings", {}), "settings") + budget_s = positive_number( + settings.get("time_limit_s", 2.0), + "settings.time_limit_s", + ) + if started_at is None: + started_at = clock() + return cls(started_at + budget_s, clock) + + def remaining_s(self, phase: str = "optimizer request") -> float: + remaining = self.expires_at - self.clock() + if remaining <= 0.0: + raise SolveDeadlineExceeded(f"{phase} deadline exceeded") + return remaining + + def check(self, phase: str = "optimizer request") -> None: + self.remaining_s(phase) diff --git a/optimizer/ftw_optimizer/direct_highs.py b/optimizer/ftw_optimizer/direct_highs.py index 64464d65..dd71c4bc 100644 --- a/optimizer/ftw_optimizer/direct_highs.py +++ b/optimizer/ftw_optimizer/direct_highs.py @@ -9,6 +9,7 @@ import numpy as np from . import SCHEMA_VERSION +from .deadline import SolveDeadline, SolveDeadlineExceeded from .model import ( _arbitrage_spread_ore_kwh, _solver_options, @@ -160,7 +161,7 @@ def solve_direct_highs( *, shared: bool = False, exact_shared_baseline: bool = False, - deadline: float | None = None, + deadline: SolveDeadline | float | None = None, prior_build_ms: float = 0.0, prior_solver_ms: float = 0.0, ) -> dict[str, Any]: @@ -171,9 +172,11 @@ def solve_direct_highs( if prepared.discrete or prepared.unsafe_cycle or prepared.unsafe_meter_split: raise DirectHighsError("direct HiGHS path requires a cycle-safe continuous tariff") if deadline is None: - deadline = started + float( - _solver_options(prepared.settings, "HIGHS")["time_limit"] + deadline = SolveDeadline( + started + + float(_solver_options(prepared.settings, "HIGHS")["time_limit"]) ) + _remaining_time_s(deadline) build_started = time.perf_counter() model = SparseModel() m = len(prepared.scenario_set.scenarios) @@ -602,8 +605,12 @@ def solve_direct_highs( time_limit_s=_remaining_time_s(deadline), ) build_ms = (time.perf_counter() - build_started) * 1000.0 + _require_ok( + highs.setOptionValue("time_limit", _remaining_time_s(deadline)), + "set service time limit", + ) solver_started = time.perf_counter() - _run_optimal(highs, "service") + _run_optimal(highs, "service", deadline) best_service = max(0.0, float(highs.getObjectiveValue())) _require_ok( highs.changeRowsBounds( @@ -623,7 +630,7 @@ def solve_direct_highs( highs.setOptionValue("time_limit", _remaining_time_s(deadline)), "set economic time limit", ) - _run_optimal(highs, "economic") + _run_optimal(highs, "economic", deadline) solver_ms = (time.perf_counter() - solver_started) * 1000.0 mip_gap = float(highs.getInfo().mip_gap) if model.integer else None solution = np.asarray(highs.getSolution().col_value, dtype=np.float64) @@ -857,10 +864,12 @@ def _add(coefficients: dict[int, float], index: int, value: float) -> None: coefficients[index] = coefficients.get(index, 0.0) + value -def _remaining_time_s(deadline: float) -> float: +def _remaining_time_s(deadline: SolveDeadline | float) -> float: + if isinstance(deadline, SolveDeadline): + return deadline.remaining_s("direct HiGHS solve") remaining = deadline - time.perf_counter() if remaining <= 0.0: - raise DirectHighsError("direct HiGHS time budget exhausted") + raise SolveDeadlineExceeded("direct HiGHS solve deadline exceeded") return remaining @@ -910,8 +919,18 @@ def _require_ok(status: highspy.HighsStatus, operation: str) -> None: raise DirectHighsError(f"HiGHS failed to {operation}: {status}") -def _run_optimal(highs: highspy.Highs, phase: str) -> None: - _require_ok(highs.run(), f"run {phase} solve") +def _run_optimal( + highs: highspy.Highs, + phase: str, + deadline: SolveDeadline | float, +) -> None: + run_status = highs.run() status = highs.getModelStatus() + if status == highspy.HighsModelStatus.kTimeLimit: + raise SolveDeadlineExceeded( + f"direct HiGHS {phase} solve deadline exceeded" + ) + _require_ok(run_status, f"run {phase} solve") if status != highspy.HighsModelStatus.kOptimal: raise DirectHighsError(f"HiGHS {phase} solve failed with status {status}") + _remaining_time_s(deadline) diff --git a/optimizer/ftw_optimizer/model.py b/optimizer/ftw_optimizer/model.py index 8ed14eaa..4373673d 100644 --- a/optimizer/ftw_optimizer/model.py +++ b/optimizer/ftw_optimizer/model.py @@ -9,6 +9,7 @@ import numpy as np from . import SCHEMA_VERSION +from .deadline import SolveDeadline, SolveDeadlineExceeded from .protocol import ProtocolError, finite_number, positive_number, require_dict, require_list @@ -244,22 +245,43 @@ def _export_price(slot: dict[str, Any], settings: dict[str, Any]) -> float: return price -def _solver_options(settings: dict[str, Any], solver: str) -> dict[str, Any]: - time_limit = positive_number(settings.get("time_limit_s", 2.0), "settings.time_limit_s") +def _solver_options( + settings: dict[str, Any], + solver: str, + deadline: SolveDeadline | None = None, +) -> dict[str, Any]: + configured_limit = positive_number( + settings.get("time_limit_s", 2.0), + "settings.time_limit_s", + ) + if deadline is None: + time_limit = max(0.05, configured_limit) + else: + time_limit = min( + configured_limit, + deadline.remaining_s(f"{solver} solve"), + ) if solver == cp.HIGHS: return { - "time_limit": max(0.05, time_limit), + "time_limit": time_limit, "mip_rel_gap": max( 0.0, finite_number(settings.get("mip_rel_gap", 0.005), "settings.mip_rel_gap"), ), } - return {"time_limit": max(0.05, time_limit)} + return {"time_limit": time_limit} -def solve(payload: dict[str, Any]) -> dict[str, Any]: +def solve( + payload: dict[str, Any], + deadline: SolveDeadline | None = None, +) -> dict[str, Any]: + started = time.perf_counter() payload = _canonicalize_storage_payload(payload) settings = require_dict(payload.get("settings", {}), "settings") + if deadline is None: + deadline = SolveDeadline.from_payload(payload, started_at=started) + deadline.check("optimizer model build") commercial = require_dict( payload.get("commercial_constraints", {}), "commercial_constraints", @@ -279,15 +301,14 @@ def solve(payload: dict[str, Any]) -> dict[str, Any]: # and thermal state can be evaluated against equally stateful telemetry. from .recourse import solve_storage_recourse - return solve_storage_recourse(payload) + return solve_storage_recourse(payload, deadline) if scenario_policy == "multistage": from .multistage import solve_storage_multistage - return solve_storage_multistage(payload) + return solve_storage_multistage(payload, deadline) if scenario_policy != "shared": raise ProtocolError("settings.scenario_policy must be shared, recourse, or multistage") - started = time.perf_counter() shared_backend = str(settings.get("shared_backend", "auto")) if shared_backend not in {"auto", "highs", "cvxpy"}: raise ProtocolError("settings.shared_backend must be auto, highs, or cvxpy") @@ -301,19 +322,23 @@ def solve(payload: dict[str, Any]) -> dict[str, Any]: from .shared_highs import DirectSharedIneligible, solve_shared_highs try: - response = solve_shared_highs(payload, started) + response = solve_shared_highs(payload, started, deadline) _validate_storage_replay( response["plan"]["actions"], require_list(payload.get("slots", []), "slots"), require_list(payload.get("storages", []), "storages"), ) return response + except SolveDeadlineExceeded: + raise except DirectSharedIneligible as exc: if shared_backend == "highs": raise ProtocolError(str(exc)) from exc + deadline.check("shared backend fallback") except Exception as exc: if shared_backend == "highs": raise + deadline.check("shared backend fallback") # The direct path is optional in auto mode. Let the reference # model validate the request again as it builds the fallback. direct_fallback_reason = str(exc) or type(exc).__name__ @@ -523,11 +548,11 @@ def solve(payload: dict[str, Any]) -> dict[str, Any]: constraints += [charge <= max_charge * direction, discharge <= max_discharge * (1 - direction)] discrete = True target = spec.get("target_energy_wh") - deadline = int(spec.get("target_slot", n - 1)) + target_slot = int(spec.get("target_slot", n - 1)) if target is not None: - deadline = min(n - 1, max(0, deadline)) + target_slot = min(n - 1, max(0, target_slot)) shortfall = cp.Variable(nonneg=True, name=f"storage_{i}_shortfall") - constraints.append(energy[deadline + 1] + shortfall >= finite_number(target, f"storages[{i}].target_energy_wh")) + constraints.append(energy[target_slot + 1] + shortfall >= finite_number(target, f"storages[{i}].target_energy_wh")) service_slack += shortfall / capacity spec["_shortfall"] = shortfall total_charge += charge @@ -629,9 +654,9 @@ def solve(payload: dict[str, Any]) -> dict[str, Any]: shortfall: cp.Variable | None = None target = spec.get("target_energy_wh") if target is not None: - deadline = min(n - 1, max(0, int(spec.get("target_slot", n - 1)))) + target_slot = min(n - 1, max(0, int(spec.get("target_slot", n - 1)))) shortfall = cp.Variable(nonneg=True, name=f"flex_{i}_shortfall") - constraints.append(energy[deadline + 1] + shortfall >= finite_number(target, f"flex_loads[{i}].target_energy_wh")) + constraints.append(energy[target_slot + 1] + shortfall >= finite_number(target, f"flex_loads[{i}].target_energy_wh")) service_slack += shortfall / capacity total_flex += power flex_loads.append(FlexVars(spec, power, energy, selection, shortfall)) @@ -931,12 +956,18 @@ def solve(payload: dict[str, Any]) -> dict[str, Any]: def run_problem(problem: cp.Problem, solver_name: str) -> None: solver = cp.HIGHS if solver_name == "HIGHS" else cp.CLARABEL - problem.solve(solver=solver, warm_start=True, **_solver_options(settings, solver)) + problem.solve( + solver=solver, + warm_start=True, + **_solver_options(settings, solver, deadline), + ) + deadline.check(f"{solver_name} solve") solver_used = preferred_solver try: run_problem(slack_problem, solver_used) except cp.error.SolverError: + deadline.check("service solver fallback") if discrete or solver_used == "CLARABEL": raise solver_used = "CLARABEL" @@ -950,6 +981,7 @@ def run_problem(problem: cp.Problem, solver_name: str) -> None: try: run_problem(cost_problem, solver_used) except cp.error.SolverError: + deadline.check("economic solver fallback") if discrete or solver_used == "CLARABEL": raise solver_used = "CLARABEL" diff --git a/optimizer/ftw_optimizer/multistage.py b/optimizer/ftw_optimizer/multistage.py index ebb374d6..c2cfd5bf 100644 --- a/optimizer/ftw_optimizer/multistage.py +++ b/optimizer/ftw_optimizer/multistage.py @@ -11,6 +11,7 @@ import numpy as np from . import SCHEMA_VERSION +from .deadline import SolveDeadline, SolveDeadlineExceeded from .model import ( OPTIMAL_STATUSES, ReplayConsistencyError, @@ -166,11 +167,18 @@ def assign(self, prepared: PreparedMultistage) -> None: ) -def solve_storage_multistage(payload: dict[str, Any]) -> dict[str, Any]: +def solve_storage_multistage( + payload: dict[str, Any], + deadline: SolveDeadline | None = None, +) -> dict[str, Any]: started = time.perf_counter() + if deadline is None: + deadline = SolveDeadline.from_payload(payload, started_at=started) + deadline.check("multistage model build") prepared_started = time.perf_counter() prepared = _prepare(payload) prepare_ms = (time.perf_counter() - prepared_started) * 1000.0 + deadline.check("multistage preparation") decomposition_threshold = _positive_int( prepared.settings.get("decomposition_threshold", 20), @@ -196,18 +204,27 @@ def solve_storage_multistage(payload: dict[str, Any]) -> dict[str, Any]: eligible, reason = ph_eligible(prepared) if eligible: try: - response = solve_progressive_hedging(prepared, started, prepare_ms) + response = solve_progressive_hedging( + prepared, + started, + prepare_ms, + deadline, + ) _validate_storage_replay( response["plan"]["actions"], prepared.slots, prepared.storages ) return response + except SolveDeadlineExceeded: + raise except ProgressiveHedgingNotConverged: if decomposition_method == "progressive_hedging": raise + deadline.check("progressive hedging fallback") decomposition = "ph-fallback-scenario-reduction-extensive-dpp" except ReplayConsistencyError as exc: if decomposition_method == "progressive_hedging": raise + deadline.check("progressive hedging replay fallback") prepared = _with_storage_discrete(prepared) decomposition = f"ph-fallback-storage-replay-{exc}" elif decomposition_method == "progressive_hedging": @@ -254,7 +271,11 @@ def solve_storage_multistage(payload: dict[str, Any]) -> dict[str, Any]: try: response = solve_direct_highs( - prepared, started, prepare_ms, decomposition.replace("-dpp", "") + prepared, + started, + prepare_ms, + decomposition.replace("-dpp", ""), + deadline=deadline, ) _validate_storage_replay( response["plan"]["actions"], prepared.slots, prepared.storages @@ -263,6 +284,7 @@ def solve_storage_multistage(payload: dict[str, Any]) -> dict[str, Any]: except (DirectHighsError, ReplayConsistencyError) as exc: if multistage_backend == "highs": raise + deadline.check("direct HiGHS fallback") direct_fallback_reason = str(exc) decomposition = f"direct-highs-fallback-{decomposition}" prepared = _with_storage_discrete(prepared) @@ -284,12 +306,23 @@ def solve_storage_multistage(payload: dict[str, Any]) -> dict[str, Any]: solver_name = "HIGHS" solver_started = time.perf_counter() try: - _run_problem(compiled.service_problem, prepared.settings, solver_name) + _run_problem( + compiled.service_problem, + prepared.settings, + solver_name, + deadline, + ) except cp.error.SolverError: + deadline.check("multistage service solver fallback") if prepared.discrete or solver_name == "CLARABEL": raise solver_name = "CLARABEL" - _run_problem(compiled.service_problem, prepared.settings, solver_name) + _run_problem( + compiled.service_problem, + prepared.settings, + solver_name, + deadline, + ) if compiled.service_problem.status not in OPTIMAL_STATUSES or compiled.service_problem.value is None: raise RuntimeError( f"multistage service-level solve failed with status {compiled.service_problem.status}" @@ -297,12 +330,23 @@ def solve_storage_multistage(payload: dict[str, Any]) -> dict[str, Any]: best_service = max(0.0, float(compiled.service_problem.value)) compiled.service_cap.value = best_service + 1e-7 try: - _run_problem(compiled.economic_problem, prepared.settings, solver_name) + _run_problem( + compiled.economic_problem, + prepared.settings, + solver_name, + deadline, + ) except cp.error.SolverError: + deadline.check("multistage economic solver fallback") if prepared.discrete or solver_name == "CLARABEL": raise solver_name = "CLARABEL" - _run_problem(compiled.economic_problem, prepared.settings, solver_name) + _run_problem( + compiled.economic_problem, + prepared.settings, + solver_name, + deadline, + ) if compiled.economic_problem.status not in OPTIMAL_STATUSES or compiled.economic_problem.value is None: raise RuntimeError( f"multistage economic solve failed with status {compiled.economic_problem.status}" @@ -328,6 +372,7 @@ def solve_storage_multistage(payload: dict[str, Any]) -> dict[str, Any]: except ReplayConsistencyError as exc: if prepared.storage_discrete: raise + deadline.check("storage replay fallback") prepared = _with_storage_discrete(prepared) direct_fallback_reason = str(exc) decomposition = f"storage-replay-fallback-{decomposition}" @@ -880,9 +925,20 @@ def _compile(prepared: PreparedMultistage, key: tuple[Any, ...]) -> CompiledMult ) -def _run_problem(problem: cp.Problem, settings: dict[str, Any], solver_name: str) -> None: +def _run_problem( + problem: cp.Problem, + settings: dict[str, Any], + solver_name: str, + deadline: SolveDeadline, +) -> None: solver = cp.HIGHS if solver_name == "HIGHS" else cp.CLARABEL - problem.solve(solver=solver, warm_start=True, enforce_dpp=True, **_solver_options(settings, solver)) + problem.solve( + solver=solver, + warm_start=True, + enforce_dpp=True, + **_solver_options(settings, solver, deadline), + ) + deadline.check(f"multistage {solver_name} solve") def _response( diff --git a/optimizer/ftw_optimizer/progressive.py b/optimizer/ftw_optimizer/progressive.py index c50cf6ac..1c31ae4a 100644 --- a/optimizer/ftw_optimizer/progressive.py +++ b/optimizer/ftw_optimizer/progressive.py @@ -9,6 +9,7 @@ import numpy as np from . import SCHEMA_VERSION +from .deadline import SolveDeadline from .model import OPTIMAL_STATUSES, _arbitrage_spread_ore_kwh, _solver_options from .protocol import ProtocolError, finite_number @@ -68,7 +69,10 @@ def ph_eligible(prepared: "PreparedMultistage") -> tuple[bool, str]: def solve_progressive_hedging( - prepared: "PreparedMultistage", started: float, prepare_ms: float + prepared: "PreparedMultistage", + started: float, + prepare_ms: float, + deadline: SolveDeadline, ) -> dict[str, Any]: eligible, reason = ph_eligible(prepared) if not eligible: @@ -82,6 +86,7 @@ def solve_progressive_hedging( _build_subproblem(prepared, si, rho_value) for si in range(len(prepared.scenario_set.scenarios)) ] + deadline.check("progressive hedging model build") build_ms = (time.perf_counter() - build_started) * 1000.0 probabilities = np.asarray([scenario.probability for scenario in prepared.scenario_set.scenarios]) @@ -94,7 +99,11 @@ def solve_progressive_hedging( subproblem.consensus_kw.value = np.zeros_like(decisions[0]) subproblem.dual_kw.value = np.zeros_like(decisions[0]) _solve_problem( - subproblem.initial_problem, settings, len(subproblems), max_iterations + subproblem.initial_problem, + settings, + len(subproblems), + max_iterations, + deadline, ) decisions = [_decision_value(subproblem) for subproblem in subproblems] consensus = _consensus_values(prepared, decisions, probabilities) @@ -106,7 +115,13 @@ def solve_progressive_hedging( for si, subproblem in enumerate(subproblems): subproblem.consensus_kw.value = consensus[si] subproblem.dual_kw.value = dual[si] - _solve_problem(subproblem.problem, settings, len(subproblems), max_iterations) + _solve_problem( + subproblem.problem, + settings, + len(subproblems), + max_iterations, + deadline, + ) decisions[si] = _decision_value(subproblem) consensus = _consensus_values(prepared, decisions, probabilities) residual_w = _nonanticipativity_residual_w(prepared, decisions, consensus) @@ -238,6 +253,7 @@ def _solve_problem( settings: dict[str, Any], scenario_count: int, max_iterations: int, + deadline: SolveDeadline, ) -> None: options_settings = dict(settings) total_limit = max(0.1, finite_number(settings.get("time_limit_s", 2), "settings.time_limit_s")) @@ -248,8 +264,9 @@ def _solve_problem( solver=cp.HIGHS, warm_start=True, enforce_dpp=True, - **_solver_options(options_settings, cp.HIGHS), + **_solver_options(options_settings, cp.HIGHS, deadline), ) + deadline.check("progressive hedging solve") if problem.status not in OPTIMAL_STATUSES or problem.value is None: raise ProgressiveHedgingNotConverged( f"PH subproblem failed with status {problem.status}" diff --git a/optimizer/ftw_optimizer/recourse.py b/optimizer/ftw_optimizer/recourse.py index bf31ecb0..ece940a2 100644 --- a/optimizer/ftw_optimizer/recourse.py +++ b/optimizer/ftw_optimizer/recourse.py @@ -9,6 +9,7 @@ import numpy as np from . import SCHEMA_VERSION +from .deadline import SolveDeadline from .model import ( OPTIMAL_STATUSES, _arbitrage_spread_ore_kwh, @@ -31,7 +32,10 @@ class ScenarioStorage: energy: cp.Variable -def solve_storage_recourse(payload: dict[str, Any]) -> dict[str, Any]: +def solve_storage_recourse( + payload: dict[str, Any], + deadline: SolveDeadline | None = None, +) -> dict[str, Any]: """Solve a two-stage stochastic storage problem. Decisions in the configured non-anticipative prefix are shared across all @@ -40,8 +44,11 @@ def solve_storage_recourse(payload: dict[str, Any]) -> dict[str, Any]: its shared first-stage action is intended for execution before replanning. """ - payload = _canonicalize_storage_payload(payload) started = time.perf_counter() + payload = _canonicalize_storage_payload(payload) + if deadline is None: + deadline = SolveDeadline.from_payload(payload, started_at=started) + deadline.check("recourse model build") settings = require_dict(payload.get("settings", {}), "settings") if require_list(payload.get("flex_loads", []), "flex_loads"): raise ProtocolError("recourse shadow does not yet support flex_loads") @@ -195,9 +202,9 @@ def solve_storage_recourse(payload: dict[str, Any]) -> dict[str, Any]: discrete = True target = spec.get("target_energy_wh") if target is not None: - deadline = min(n - 1, max(0, int(spec.get("target_slot", n - 1)))) + target_slot = min(n - 1, max(0, int(spec.get("target_slot", n - 1)))) shortfall = cp.Variable(nonneg=True, name=f"scenario_{si}_storage_{i}_shortfall") - constraints.append(energy[deadline + 1] + shortfall >= finite_number(target, f"storages[{i}].target_energy_wh")) + constraints.append(energy[target_slot + 1] + shortfall >= finite_number(target, f"storages[{i}].target_energy_wh")) scenario_service += shortfall / capacity cycle_ore = max(0.0, finite_number(spec.get("cycle_cost_ore_kwh", 0), "storage.cycle_cost_ore_kwh")) @@ -302,13 +309,19 @@ def solve_storage_recourse(payload: dict[str, Any]) -> dict[str, Any]: def run_problem(problem: cp.Problem, solver_name: str) -> None: solver = cp.HIGHS if solver_name == "HIGHS" else cp.CLARABEL - problem.solve(solver=solver, warm_start=True, **_solver_options(settings, solver)) + problem.solve( + solver=solver, + warm_start=True, + **_solver_options(settings, solver, deadline), + ) + deadline.check(f"recourse {solver_name} solve") slack_problem = cp.Problem(cp.Minimize(worst_service_slack), constraints) solver_used = preferred_solver try: run_problem(slack_problem, solver_used) except cp.error.SolverError: + deadline.check("recourse service solver fallback") if discrete or solver_used == "CLARABEL": raise solver_used = "CLARABEL" @@ -323,6 +336,7 @@ def run_problem(problem: cp.Problem, solver_name: str) -> None: try: run_problem(cost_problem, solver_used) except cp.error.SolverError: + deadline.check("recourse economic solver fallback") if discrete or solver_used == "CLARABEL": raise solver_used = "CLARABEL" diff --git a/optimizer/ftw_optimizer/shared_highs.py b/optimizer/ftw_optimizer/shared_highs.py index ea5c5fbc..fcd54059 100644 --- a/optimizer/ftw_optimizer/shared_highs.py +++ b/optimizer/ftw_optimizer/shared_highs.py @@ -11,7 +11,8 @@ SharedBaselineReplayError, solve_direct_highs, ) -from .model import _STORAGE_INITIAL_ABOVE_MAXIMUM_KEY, _solver_options +from .deadline import SolveDeadline +from .model import _STORAGE_INITIAL_ABOVE_MAXIMUM_KEY from .multistage import _prepare from .protocol import ProtocolError, finite_number, require_dict, require_list from .scenario_tree import ScenarioTree @@ -21,7 +22,11 @@ class DirectSharedIneligible(RuntimeError): pass -def solve_shared_highs(payload: dict[str, Any], started: float) -> dict[str, Any]: +def solve_shared_highs( + payload: dict[str, Any], + started: float, + deadline: SolveDeadline, +) -> dict[str, Any]: """Solve shared storage through the sparse HiGHS builder.""" prepared_started = time.perf_counter() direct_payload, risk_alpha = _direct_payload(payload) @@ -119,9 +124,6 @@ def solve_shared_highs(payload: dict[str, Any], started: float) -> dict[str, Any ), ) prepare_ms = (time.perf_counter() - prepared_started) * 1000.0 - deadline = started + float( - _solver_options(prepared.settings, "HIGHS")["time_limit"] - ) try: return solve_direct_highs( prepared, @@ -132,6 +134,7 @@ def solve_shared_highs(payload: dict[str, Any], started: float) -> dict[str, Any deadline=deadline, ) except SharedBaselineReplayError as exc: + deadline.check("shared baseline retry") return solve_direct_highs( prepared, started, diff --git a/optimizer/ftw_optimizer/worker.py b/optimizer/ftw_optimizer/worker.py index 4aed8a69..725dc0b6 100644 --- a/optimizer/ftw_optimizer/worker.py +++ b/optimizer/ftw_optimizer/worker.py @@ -1,22 +1,25 @@ from __future__ import annotations +import argparse import ctypes import gc -import argparse import importlib.metadata import json import os import socket import sys import threading +import time import traceback +from collections.abc import Callable from pathlib import Path from typing import Any import cvxpy as cp +from .deadline import SolveDeadline, SolveDeadlineExceeded from .model import solve -from .protocol import ProtocolError, error_response, parse_request +from .protocol import ParsedRequest, ProtocolError, error_response, parse_request # PROTOCOL_VERSION is what this worker speaks; MIN_PROTOCOL_VERSION is the @@ -45,14 +48,35 @@ def release_unused_memory() -> None: malloc_trim(0) -def handle(raw: Any) -> dict[str, Any]: +def handle( + raw: Any, + *, + received_at: float | None = None, + clock: Callable[[], float] = time.perf_counter, + parsed: ParsedRequest | None = None, + deadline: SolveDeadline | None = None, +) -> dict[str, Any]: + if received_at is None: + received_at = clock() request_id = "unknown" try: - parsed = parse_request(raw) + if parsed is None: + parsed = parse_request(raw) request_id = parsed.request_id - return solve(parsed.payload) + if deadline is None: + deadline = SolveDeadline.from_payload( + parsed.payload, + started_at=received_at, + clock=clock, + ) + deadline.check("optimizer queue") + response = solve(parsed.payload, deadline=deadline) + deadline.check("optimizer response") + return response except ProtocolError as exc: return error_response(request_id, "invalid_request", str(exc)) + except SolveDeadlineExceeded as exc: + return error_response(request_id, "deadline_exceeded", str(exc)) except cp.error.SolverError as exc: return error_response(request_id, "solver_error", str(exc)) except Exception as exc: # worker boundary: one bad request must not kill the process @@ -78,10 +102,16 @@ def handshake(raw: Any) -> dict[str, Any] | None: } -def process_stream(reader: Any, writer: Any) -> None: +def process_stream( + reader: Any, + writer: Any, + *, + clock: Callable[[], float] = time.perf_counter, +) -> None: for line in reader: if not line.strip(): continue + received_at = clock() try: raw = json.loads(line) except json.JSONDecodeError as exc: @@ -91,9 +121,34 @@ def process_stream(reader: Any, writer: Any) -> None: if response is None: # Handshakes stay responsive while a solve is in progress, # but solver state and its memory cleanup remain serialized. - with SOLVE_LOCK: - response = handle(raw) - try: + request_id = "unknown" + 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): + response = error_response( + parsed.request_id, + "deadline_exceeded", + "optimizer queue deadline exceeded", + ) writer.write( json.dumps( response, @@ -103,10 +158,31 @@ def process_stream(reader: Any, writer: Any) -> None: + "\n" ) writer.flush() + continue + try: + response = handle( + raw, + received_at=received_at, + clock=clock, + parsed=parsed, + deadline=deadline, + ) + try: + writer.write( + json.dumps( + response, + separators=(",", ":"), + allow_nan=False, + ) + + "\n" + ) + writer.flush() + finally: + response = None + release_unused_memory() finally: - response = None - release_unused_memory() - continue + SOLVE_LOCK.release() + 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 new file mode 100644 index 00000000..c8f95fb7 --- /dev/null +++ b/optimizer/tests/test_deadline.py @@ -0,0 +1,117 @@ +from __future__ import annotations + +import cvxpy as cp +import highspy +import pytest + +from ftw_optimizer import shared_highs +from ftw_optimizer.deadline import SolveDeadline, SolveDeadlineExceeded +from ftw_optimizer.direct_highs import ( + DirectHighsError, + _remaining_time_s, + _run_optimal, +) +from ftw_optimizer.model import _solver_options, solve + + +class FakeClock: + def __init__(self, now: float = 0.0) -> None: + self.now = now + + def __call__(self) -> float: + return self.now + + def advance(self, seconds: float) -> None: + self.now += seconds + + +def test_one_deadline_shrinks_across_cvxpy_and_direct_highs_phases() -> None: + clock = FakeClock(10.0) + deadline = SolveDeadline.from_payload( + {"settings": {"time_limit_s": 1.0}}, + started_at=clock(), + clock=clock, + ) + settings = {"time_limit_s": 1.0} + + assert _solver_options(settings, cp.HIGHS, deadline)["time_limit"] == pytest.approx(1.0) + clock.advance(0.6) + assert _solver_options(settings, cp.CLARABEL, deadline)["time_limit"] == pytest.approx(0.4) + assert _remaining_time_s(deadline) == pytest.approx(0.4) + + # The old per-solve 50 ms floor must not extend the request deadline. + clock.advance(0.39) + assert _solver_options(settings, cp.HIGHS, deadline)["time_limit"] == pytest.approx(0.01) + clock.advance(0.02) + with pytest.raises(SolveDeadlineExceeded, match="deadline exceeded"): + _solver_options(settings, cp.HIGHS, deadline) + with pytest.raises(SolveDeadlineExceeded, match="deadline exceeded"): + _remaining_time_s(deadline) + + +def test_deadline_error_bypasses_shared_backend_fallback(monkeypatch) -> None: + deadline = SolveDeadline(1.0, FakeClock()) + direct_calls: list[SolveDeadline] = [] + + def fail_direct( + _payload: dict, + _started: float, + received_deadline: SolveDeadline, + ) -> dict: + direct_calls.append(received_deadline) + raise SolveDeadlineExceeded("direct HiGHS solve deadline exceeded") + + monkeypatch.setattr(shared_highs, "solve_shared_highs", fail_direct) + + with pytest.raises(SolveDeadlineExceeded, match="deadline exceeded"): + solve( + { + "settings": { + "shared_backend": "auto", + "time_limit_s": 10.0, + }, + "commercial_constraints": {}, + "slots": [{}], + "storages": [], + }, + deadline=deadline, + ) + + assert direct_calls == [deadline] + + +class FakeHighs: + def __init__( + self, + status: highspy.HighsModelStatus, + run_status: highspy.HighsStatus = highspy.HighsStatus.kOk, + ) -> None: + self.status = status + self.run_status = run_status + + def run(self) -> highspy.HighsStatus: + return self.run_status + + def getModelStatus(self) -> highspy.HighsModelStatus: + return self.status + + +def test_direct_highs_time_limit_is_a_deadline_not_a_fallback_error() -> None: + deadline = SolveDeadline(1.0, FakeClock()) + + with pytest.raises(SolveDeadlineExceeded, match="service solve deadline exceeded"): + _run_optimal( + FakeHighs( + highspy.HighsModelStatus.kTimeLimit, + highspy.HighsStatus.kWarning, + ), + "service", + deadline, + ) + + with pytest.raises(DirectHighsError, match="failed with status"): + _run_optimal( + FakeHighs(highspy.HighsModelStatus.kInfeasible), + "service", + deadline, + ) diff --git a/optimizer/tests/test_model.py b/optimizer/tests/test_model.py index 5600c7ac..a18e565a 100644 --- a/optimizer/tests/test_model.py +++ b/optimizer/tests/test_model.py @@ -10,6 +10,7 @@ import numpy as np import pytest +from ftw_optimizer.deadline import SolveDeadlineExceeded from ftw_optimizer.direct_highs import DirectHighsError, _remaining_time_s from ftw_optimizer.multistage import clear_multistage_cache from ftw_optimizer.model import ( @@ -264,7 +265,7 @@ def test_direct_highs_accepts_a_positive_sub_50ms_budget() -> None: def test_direct_highs_rejects_an_exhausted_budget() -> None: - with pytest.raises(DirectHighsError, match="time budget exhausted"): + with pytest.raises(SolveDeadlineExceeded, match="deadline exceeded"): _remaining_time_s(time.perf_counter() - 0.001) diff --git a/optimizer/tests/test_worker.py b/optimizer/tests/test_worker.py index 59b31246..abbcdbde 100644 --- a/optimizer/tests/test_worker.py +++ b/optimizer/tests/test_worker.py @@ -1,6 +1,7 @@ from __future__ import annotations import io +import json import threading from ftw_optimizer import worker @@ -15,7 +16,7 @@ def test_health_stays_responsive_without_cleaning_memory_during_solve( thread_errors: list[BaseException] = [] monkeypatch.setattr(worker, "SOLVE_LOCK", threading.Lock()) - def fake_handle(_raw: object) -> dict[str, object]: + def fake_handle(_raw: object, **_kwargs: object) -> dict[str, object]: solve_started.set() if not finish_solve.wait(timeout=2): raise TimeoutError("test did not release solve") @@ -63,3 +64,153 @@ def run_solve() -> None: assert not solve_thread.is_alive() assert thread_errors == [] assert cleanup_calls == [True] + + +class FakeClock: + def __init__(self) -> None: + self.value = 0.0 + self.condition = threading.Condition() + + def __call__(self) -> float: + with self.condition: + return self.value + + def advance(self, seconds: float) -> None: + with self.condition: + self.value += seconds + self.condition.notify_all() + + +class FakeSolveLock: + def __init__(self, clock: FakeClock) -> None: + self.clock = clock + self.held = False + self.queued = threading.Event() + + def acquire(self, blocking: bool = True, timeout: float = -1) -> bool: + with self.clock.condition: + if not self.held: + self.held = True + return True + if not blocking: + return False + self.queued.set() + expires_at = self.clock.value + timeout + while self.held: + if timeout >= 0 and self.clock.value >= expires_at: + return False + self.clock.condition.wait() + self.held = True + return True + + def release(self) -> None: + with self.clock.condition: + self.held = False + self.clock.condition.notify_all() + + def locked(self) -> bool: + with self.clock.condition: + return self.held + + +def test_expired_request_leaves_solve_queue_without_running( + monkeypatch, +) -> None: + clock = FakeClock() + solve_lock = FakeSolveLock(clock) + first_started = threading.Event() + finish_first = 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 == "first": + first_started.set() + if not finish_first.wait(timeout=2): + raise TimeoutError("test did not release first solve") + return {"ok": True, "request_id": request_id} + + monkeypatch.setattr(worker, "solve", fake_solve) + + def request(request_id: str, budget_s: float) -> io.StringIO: + return io.StringIO( + json.dumps( + { + "schema_version": 1, + "request_id": request_id, + "settings": {"time_limit_s": budget_s}, + "slots": [{}], + } + ) + + "\n" + ) + + def run(request_id: str, budget_s: float, output: io.StringIO) -> None: + try: + worker.process_stream( + request(request_id, budget_s), + output, + clock=clock, + ) + except BaseException as exc: + thread_errors.append(exc) + + first_output = io.StringIO() + first = threading.Thread(target=run, args=("first", 10.0, first_output)) + first.start() + assert first_started.wait(timeout=1) + + expired_output = io.StringIO() + expired = threading.Thread(target=run, args=("expired", 1.0, expired_output)) + expired.start() + assert solve_lock.queued.wait(timeout=1) + clock.advance(2.0) + expired.join(timeout=1) + assert not expired.is_alive() + assert solve_lock.locked() + assert json.loads(expired_output.getvalue())["error"]["code"] == "deadline_exceeded" + assert solve_calls == ["first"] + + finish_first.set() + first.join(timeout=1) + assert not first.is_alive() + + fresh_output = io.StringIO() + fresh = threading.Thread(target=run, args=("fresh", 1.0, fresh_output)) + fresh.start() + fresh.join(timeout=1) + assert not fresh.is_alive() + assert json.loads(fresh_output.getvalue())["ok"] is True + assert solve_calls == ["first", "fresh"] + assert thread_errors == [] + + +def test_handle_rejects_a_result_that_finishes_after_its_deadline( + monkeypatch, +) -> None: + clock = FakeClock() + solve_calls = 0 + + def fake_solve(_payload: dict, **_kwargs: object) -> dict[str, object]: + nonlocal solve_calls + solve_calls += 1 + clock.advance(2.0) + return {"ok": True} + + monkeypatch.setattr(worker, "solve", fake_solve) + response = worker.handle( + { + "schema_version": 1, + "request_id": "late", + "settings": {"time_limit_s": 1.0}, + "slots": [{}], + }, + clock=clock, + ) + + assert solve_calls == 1 + assert response["error"]["code"] == "deadline_exceeded"