diff --git a/CHANGELOG.md b/CHANGELOG.md index e94db42..6436265 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,61 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Features +- Servers now isolate each client's state in a **job**, so a single discipline + server can support concurrent clients. Previously a server held one + discipline instance and shared it: a second client's `Setup` ran + `_clear_data()` and rebuilt `_var_meta` while the first client was still + using it, and `SetOptions` and the in-place shape edits of + `SetVariableShapes` collided the same way. The victim either aborted with + `INTERNAL` on a `KeyError` or, worse, had a short array read into a longer + buffer and received a **zero-padded result with no error** -- which an + optimizer would happily consume. `examples/rosenbrock.py` was the sharp + case, since its variable shape derives from an option, so whichever client + called `SetOptions` last fixed the shapes both clients got. +- A job is a session owning one discipline instance, so the state an author + keeps on `self` -- a mesh, a solver, an `om.Problem` -- is private to one + client. **Every existing discipline hook signature is unchanged**; only the + server constructor moves, from `ExplicitServer(discipline=Paraboloid())` to + `ExplicitServer(discipline=Paraboloid)`. A class works directly as a + factory when its `initialize()` does its own configuration; a discipline + configured externally needs a closure or `functools.partial`. `Discipline` + gains a `job` attribute, giving `self.job.job_id`, and an optional + `teardown_job()` hook called when a job ends or is evicted. +- Three RPCs are added to `DisciplineService`: `StartJob`, `EndJob` and + `KeepAlive`. The job id travels in a `philote-job-id` metadata header rather + than a message field, which costs 7.7 us on a unary call and 13.9 us on a + stream; HPACK indexes the repeated value, so after the first call it is a + byte or two on the wire. A field would instead have been re-serialized on + every chunk of every array with only the first ever read, and each of the + unary RPCs -- all of which take `google.protobuf.Empty` -- would have needed + its own request message. Clients attach the header through a channel + interceptor, so no call site passes it. `GetInfo` and `GetAvailableOptions` + stay job-independent, since they describe the discipline class. +- Clients start a job lazily on the first call that needs one, so existing + scripts and both OpenMDAO components work unchanged. `start_job()`, + `end_job()`, `keep_alive()` and a `job()` context manager are available for + explicit control. An unknown or expired job raises the new `PhiloteJobError` + and is **not** silently replaced: the state that job held is gone, and an + optimizer continuing against a fresh discipline would return plausible but + wrong results. As part of this, every client method now translates gRPC + errors into Philote exceptions; previously only the compute calls did, and a + server-side failure during `run_setup` escaped as a raw `grpc.RpcError`. +- Servers cap concurrent jobs (`max_jobs`, default 8) and evict idle ones + (`ttl`, default one hour), because a job can hold a mesh or a live solver and + a client that dies would otherwise leak it. Exceeding the cap returns + `RESOURCE_EXHAUSTED` rather than exhausting memory. Note that the gRPC + thread pool, not `max_jobs`, is the cap that actually binds -- every in-flight + RPC holds a worker for its whole duration -- so the server warns at startup + when the pool is smaller than the job limit. +- Separate jobs may evaluate concurrently, with no global lock in the path, but + the GIL decides whether that yields throughput. Measured with four + concurrent clients: a pure-Python discipline sees 1.0x (278 ms to 1107 ms), a + NumPy `A @ A` sees 0.8x because threaded BLAS already saturates the cores + with one call, and a discipline whose compiled solver releases the GIL sees + 4.0x (306 ms to 310 ms). Jobs buy correctness unconditionally and throughput + conditionally; pure-Python disciplines, including `OpenMdaoSubProblem`, become + correct under concurrent clients rather than faster. + - Continuous array data is now read and written through the packed wire buffer directly, rather than through the protobuf `repeated double` API, which converts every element to and from a boxed Python float. A packed diff --git a/docs/docs/getting-started/quickstart.md b/docs/docs/getting-started/quickstart.md index aae3fb7..b133dbb 100644 --- a/docs/docs/getting-started/quickstart.md +++ b/docs/docs/getting-started/quickstart.md @@ -68,20 +68,35 @@ from concurrent import futures import grpc # ... -server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) +server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) ``` -Next, the **Paraboloid** discipline is attached to the server: +Next, the **Paraboloid** discipline is attached to the server. Note the class +`Paraboloid`, not an instance `Paraboloid()` — the server builds one discipline +per job, which is what lets several clients share a server without overwriting +each other's setup, so it needs something it can call: ```python import philote_mdo.general as pmdo from philote_mdo.examples import Paraboloid # ... -discipline = pmdo.ExplicitServer(discipline=Paraboloid()) +discipline = pmdo.ExplicitServer(discipline=Paraboloid) discipline.attach_to_server(server) ``` +A class works directly as a factory when its `initialize()` performs its own +configuration. A discipline that has to be configured from the outside needs a +closure or `functools.partial` instead: + +```python +from functools import partial + +discipline = pmdo.ExplicitServer( + discipline=partial(MyDiscipline, mesh_file="wing.cgns") +) +``` + Finally, the port of the server is defined (opening a port is necessary for network communication) and the server is started: ```python @@ -106,9 +121,9 @@ import philote_mdo.general as pmdo from philote_mdo.examples import Paraboloid -server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) +server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) -discipline = pmdo.ExplicitServer(discipline=Paraboloid()) +discipline = pmdo.ExplicitServer(discipline=Paraboloid) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -117,6 +132,46 @@ print("Server started. Listening on port 50051.") server.wait_for_termination() ``` +## Jobs + +Each client gets its own **job** on the server: a session owning one discipline +instance and everything built on it, including the options it set and the +variable metadata `Setup` produced. Two clients can therefore use one server at +the same time without interfering. + +The client starts a job on its first call that needs one, so the code above +needs no changes to benefit. When you want to control the lifetime explicitly, +use the context manager, which releases the server's resources on exit: + +```python +with client.job(): + client.run_setup() + client.get_variable_definitions() + outputs = client.run_compute(inputs) +``` + +If the server no longer recognises a job -- it expired, or the server was +restarted -- the client raises `PhiloteJobError`. It deliberately does not +start a replacement job on your behalf, because whatever state the old job held +is gone, and an optimizer that carried on would return plausible but wrong +results. + +:::note +Servers cap the number of concurrent jobs (`max_jobs`) and evict jobs that go +unused (`ttl`). Make sure the gRPC thread pool has at least as many workers as +the job cap: every in-flight RPC holds a worker for its whole duration, so a +pool smaller than the cap makes jobs queue instead of run. The server warns at +startup when the two are mismatched. +::: + +:::warning +Separate jobs may evaluate at the same time, but Python's GIL decides whether +that produces a speedup. A discipline wrapping a compiled solver releases the +GIL and genuinely runs in parallel. A pure-Python discipline does not: jobs make +it *correct* under concurrent clients, not faster. For parallel evaluation of +pure-Python disciplines, run several server processes. +::: + ## Calling the Discipline Using a Client Now that a server is running, it can be queried using a client. Philote-Python offers a number of clients for this purpose, ranging from the general implementation to OpenMDAO and CSDL components. However, under the hood, the OpenMDAO and CSDL components use the general client implementation. diff --git a/docs/docs/tutorials/implicit-disciplines.md b/docs/docs/tutorials/implicit-disciplines.md index 6a4ec07..7fc7396 100644 --- a/docs/docs/tutorials/implicit-disciplines.md +++ b/docs/docs/tutorials/implicit-disciplines.md @@ -227,14 +227,13 @@ import grpc import philote_mdo.general as pmdo def run_server(): - # Create the discipline - discipline = QuadraticSolver() - - # Create gRPC server - server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) + # Create gRPC server. The pool should be at least as large as the + # server's job cap, since every in-flight RPC holds a worker. + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) - # Create and attach implicit server - impl_server = pmdo.ImplicitServer(discipline=discipline) + # Create and attach implicit server. It takes a factory and builds one + # discipline per job, so concurrent clients do not interfere. + impl_server = pmdo.ImplicitServer(discipline=QuadraticSolver) impl_server.attach_to_server(server) # Start server @@ -268,7 +267,7 @@ def run_production_server(): ] ) - impl_server = pmdo.ImplicitServer(discipline=discipline) + impl_server = pmdo.ImplicitServer(discipline=QuadraticSolver) impl_server.attach_to_server(server) # Use secure connection in production diff --git a/examples/openmdao_sellar.py b/examples/openmdao_sellar.py index d0a0253..0ff09ba 100644 --- a/examples/openmdao_sellar.py +++ b/examples/openmdao_sellar.py @@ -34,9 +34,15 @@ def run(): - server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) + # the pool must be at least as large as the server's job cap: every + # in-flight RPC holds a worker for its whole duration + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) - discipline = pmdo.ExplicitServer(discipline=SellarGroup()) + # One discipline instance is built per job, so concurrent clients do not + # interfere. A class works as the factory here because its initialize() + # does its own configuration; a discipline configured from the outside + # needs a closure or functools.partial. + discipline = pmdo.ExplicitServer(discipline=SellarGroup) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") diff --git a/examples/parabaloid_explicit.py b/examples/parabaloid_explicit.py index f9dd210..dd1a0a6 100644 --- a/examples/parabaloid_explicit.py +++ b/examples/parabaloid_explicit.py @@ -34,9 +34,15 @@ def run(): - server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) + # the pool must be at least as large as the server's job cap: every + # in-flight RPC holds a worker for its whole duration + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) - discipline = pmdo.ExplicitServer(discipline=Paraboloid()) + # One discipline instance is built per job, so concurrent clients do not + # interfere. A class works as the factory here because its initialize() + # does its own configuration; a discipline configured from the outside + # needs a closure or functools.partial. + discipline = pmdo.ExplicitServer(discipline=Paraboloid) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") diff --git a/examples/quadratic_implicit.py b/examples/quadratic_implicit.py index d1818d4..d73e73c 100644 --- a/examples/quadratic_implicit.py +++ b/examples/quadratic_implicit.py @@ -34,8 +34,14 @@ def run(): - server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = pmdo.ImplicitServer(discipline=QuadradicImplicit()) + # the pool must be at least as large as the server's job cap: every + # in-flight RPC holds a worker for its whole duration + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) + # One discipline instance is built per job, so concurrent clients do not + # interfere. A class works as the factory here because its initialize() + # does its own configuration; a discipline configured from the outside + # needs a closure or functools.partial. + discipline = pmdo.ImplicitServer(discipline=QuadradicImplicit) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") diff --git a/philote_mdo/general/__init__.py b/philote_mdo/general/__init__.py index e5ef971..ebe5a34 100644 --- a/philote_mdo/general/__init__.py +++ b/philote_mdo/general/__init__.py @@ -35,6 +35,8 @@ from .explicit_server import ExplicitServer from .implicit_server import ImplicitServer +from .job import JOB_METADATA_KEY, Job, JobState, JobStore + from .discipline import Discipline from .explicit_discipline import ExplicitDiscipline from .implicit_discipline import ImplicitDiscipline diff --git a/philote_mdo/general/discipline.py b/philote_mdo/general/discipline.py index e0baaa5..c9a47c5 100644 --- a/philote_mdo/general/discipline.py +++ b/philote_mdo/general/discipline.py @@ -67,6 +67,11 @@ def __init__(self): # flag that indicates the discipline is implicit self._is_implicit = False + # the job that owns this instance, assigned by the server when the + # job is created. One discipline instance serves exactly one job, so + # anything stored on self is private to that client. + self.job = None + # dictionary of available discipline options (with types) self.options_list = {} @@ -283,6 +288,17 @@ def setup_partials(self): def configure(self): pass + def teardown_job(self): + """ + Releases whatever this instance holds, before its job is discarded. + + Called when the client ends the job and when the server evicts one + that has gone idle, so a discipline that opened a file, started a + subprocess, or built a solver should close it here. Overriding this + is optional; the default does nothing. + """ + pass + def _clear_data(self): """ Clears all metadata of the discipline. diff --git a/philote_mdo/general/discipline_client.py b/philote_mdo/general/discipline_client.py index 045197e..a7a5528 100644 --- a/philote_mdo/general/discipline_client.py +++ b/philote_mdo/general/discipline_client.py @@ -27,19 +27,106 @@ # the linked websites, of the information, products, or services contained # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. +import contextlib + +import grpc import numpy as np import google.protobuf.empty_pb2 as empty import philote_mdo.generated.data_pb2 as data import philote_mdo.generated.disciplines_pb2_grpc as disc import philote_mdo.utils as utils from philote_mdo.general.discipline_server import _python_to_value, _value_to_python +from philote_mdo.general.job import JOB_METADATA_KEY from philote_mdo.utils.validation import ( + JobCapacityError, + JobNotFoundError, + PhiloteServerError, PhiloteValidationError, validate_is_dict, validate_numpy_array, ) +class _JobMetadataInterceptor( + grpc.UnaryUnaryClientInterceptor, + grpc.UnaryStreamClientInterceptor, + grpc.StreamUnaryClientInterceptor, + grpc.StreamStreamClientInterceptor, +): + """ + Attaches the client's current job id to every outgoing call. + + Sitting under the stubs means no call site has to pass the header itself, + and a job started part way through a session is picked up from the next + call onwards. HPACK indexes the repeated value, so after the first call the + id costs a byte or two on the wire. + """ + + def __init__(self, client): + self._client = client + + def _augment(self, client_call_details): + job_id = self._client._job_id + + if job_id is None: + return client_call_details + + metadata = list(client_call_details.metadata or []) + metadata.append((JOB_METADATA_KEY, job_id)) + + return client_call_details._replace(metadata=metadata) + + def intercept_unary_unary(self, continuation, details, request): + return continuation(self._augment(details), request) + + def intercept_unary_stream(self, continuation, details, request): + return continuation(self._augment(details), request) + + def intercept_stream_unary(self, continuation, details, request_iterator): + return continuation(self._augment(details), request_iterator) + + def intercept_stream_stream(self, continuation, details, request_iterator): + return continuation(self._augment(details), request_iterator) + + +def raise_for_rpc_error(error, context): + """ + Re-raises a gRPC error as the matching Philote exception. + + Parameters + ---------- + error : grpc.RpcError + The error raised by the stub call. + context : str + Name of the client method, used in the message. + + Raises + ------ + JobNotFoundError + If the server did not recognise the job. This is terminal: the state + the job held is gone, so a caller must start a new job and set it up + again rather than carry on. + JobCapacityError + If the server is already holding as many jobs as it allows. Unlike the + above this is worth retrying, once another client has finished. + PhiloteServerError + For every other server-side failure. + """ + if error.code() == grpc.StatusCode.NOT_FOUND: + raise JobNotFoundError( + f"Server rejected the job during {context}: {error.details()}" + ) from error + + if error.code() == grpc.StatusCode.RESOURCE_EXHAUSTED: + raise JobCapacityError( + f"Server is at capacity during {context}: {error.details()}" + ) from error + + raise PhiloteServerError( + f"Server error during {context}: {error.details()}" + ) from error + + class DisciplineClient: """ Base class for analysis discipline clients. @@ -52,6 +139,16 @@ def __init__(self, channel): # grpc options self.grpc_options = [] + # server-side job this client is bound to. Started lazily on the first + # call that needs one, so existing scripts need no changes. + self._job_id = None + + # every stub is built on the intercepted channel, so the job header + # rides along without any call site passing it + self._channel = grpc.intercept_channel( + channel, _JobMetadataInterceptor(self) + ) + # discipline properties self._name = "" self._version = "" @@ -60,7 +157,7 @@ def __init__(self, channel): self._provides_gradients = False # discipline client stub - self._disc_stub = disc.DisciplineServiceStub(channel) + self._disc_stub = disc.DisciplineServiceStub(self._channel) # streaming options # doubles per message. The cost of a stream is dominated by the number @@ -78,11 +175,101 @@ def __init__(self, channel): # list of available options self.options_list = {} + @property + def job_id(self): + """ + The server-side job this client is bound to, or None. + """ + return self._job_id + + def start_job(self): + """ + Starts a job on the server and binds this client to it. + + Called automatically by the first RPC that needs a job, so most code + never has to call it. Call it directly when the job id is wanted up + front, or to control exactly when server-side resources are claimed. + + Returns + ------- + str + The new job id. + """ + if self._job_id is not None: + return self._job_id + + try: + handle = self._disc_stub.StartJob(empty.Empty()) + except grpc.RpcError as e: + raise_for_rpc_error(e, "start_job") + + self._job_id = handle.job_id + + return self._job_id + + def end_job(self): + """ + Ends this client's job, releasing what the server holds for it. + + Does nothing when no job has been started. + """ + if self._job_id is None: + return + + try: + self._disc_stub.EndJob(empty.Empty()) + except grpc.RpcError as e: + raise_for_rpc_error(e, "end_job") + finally: + self._job_id = None + + def keep_alive(self): + """ + Tells the server this job is still wanted. + + Useful when an optimizer sits idle between design iterations for + longer than the server's job time-to-live. + """ + self._ensure_job() + + try: + self._disc_stub.KeepAlive(empty.Empty()) + except grpc.RpcError as e: + raise_for_rpc_error(e, "keep_alive") + + @contextlib.contextmanager + def job(self): + """ + Runs a block against a job, ending it afterwards. + + Examples + -------- + >>> with client.job(): + ... client.run_setup() + ... outputs = client.run_compute(inputs) + """ + self.start_job() + + try: + yield self + finally: + self.end_job() + + def _ensure_job(self): + """ + Starts a job if this client has not got one yet. + """ + if self._job_id is None: + self.start_job() + def get_discipline_info(self): """ Gets the discipline properties from the analysis server. """ - response = self._disc_stub.GetInfo(empty.Empty()) + try: + response = self._disc_stub.GetInfo(empty.Empty()) + except grpc.RpcError as e: + raise_for_rpc_error(e, "get_discipline_info") self._is_continuous = response.continuous self._is_differentiable = response.differentiable self._provides_gradients = response.provides_gradients @@ -93,13 +280,21 @@ def send_stream_options(self): """ Transmits the stream options for the remote analysis to the server. """ - self._disc_stub.SetStreamOptions(self._stream_options) + self._ensure_job() + + try: + self._disc_stub.SetStreamOptions(self._stream_options) + except grpc.RpcError as e: + raise_for_rpc_error(e, "send_stream_options") def get_available_options(self): """ Gets the available options for the analysis discipline. """ - opts = self._disc_stub.GetAvailableOptions(empty.Empty()) + try: + opts = self._disc_stub.GetAvailableOptions(empty.Empty()) + except grpc.RpcError as e: + raise_for_rpc_error(e, "get_available_options") for name, val in zip(opts.options, opts.type): type_str = None @@ -129,15 +324,24 @@ def send_options(self, options): None """ validate_is_dict(options, "send_options") + self._ensure_job() proto_options = data.DisciplineOptions() proto_options.options.update(options) - self._disc_stub.SetOptions(proto_options) + try: + self._disc_stub.SetOptions(proto_options) + except grpc.RpcError as e: + raise_for_rpc_error(e, "send_options") def run_setup(self): """ Runs the setup function on the analysis server. """ - self._disc_stub.Setup(empty.Empty()) + self._ensure_job() + + try: + self._disc_stub.Setup(empty.Empty()) + except grpc.RpcError as e: + raise_for_rpc_error(e, "run_setup") def get_variable_definitions(self): """ @@ -151,17 +355,22 @@ def get_variable_definitions(self): client is reused across jobs) mirrors the server, which clears its metadata at the start of every ``Setup``. """ + self._ensure_job() + self._var_meta = [] self._discrete_var_meta = [] - for message in self._disc_stub.GetVariableDefinitions(empty.Empty()): - if message.type in ( - data.VariableType.kDiscreteInput, - data.VariableType.kDiscreteOutput, - ): - self._discrete_var_meta += [message] - else: - self._var_meta += [message] + try: + for message in self._disc_stub.GetVariableDefinitions(empty.Empty()): + if message.type in ( + data.VariableType.kDiscreteInput, + data.VariableType.kDiscreteOutput, + ): + self._discrete_var_meta += [message] + else: + self._var_meta += [message] + except grpc.RpcError as e: + raise_for_rpc_error(e, "get_variable_definitions") def get_partials_definitions(self): """ @@ -171,10 +380,15 @@ def get_partials_definitions(self): rather than appended to, so repeated calls do not accumulate duplicate partials entries. """ + self._ensure_job() + self._partials_meta = [] - for message in self._disc_stub.GetPartialDefinitions(empty.Empty()): - self._partials_meta += [message] + try: + for message in self._disc_stub.GetPartialDefinitions(empty.Empty()): + self._partials_meta += [message] + except grpc.RpcError as e: + raise_for_rpc_error(e, "get_partials_definitions") def get_dynamic_variables(self): """ @@ -220,7 +434,12 @@ def send_variable_shapes(self, variable_metadata): variable_metadata : list of VariableMetaData shapes for dynamic variables """ - self._disc_stub.SetVariableShapes(iter(variable_metadata)) + self._ensure_job() + + try: + self._disc_stub.SetVariableShapes(iter(variable_metadata)) + except grpc.RpcError as e: + raise_for_rpc_error(e, "send_variable_shapes") # index by type and name once; searching the list per shape is # quadratic in the number of variables diff --git a/philote_mdo/general/discipline_server.py b/philote_mdo/general/discipline_server.py index 1dff413..ab6ea90 100644 --- a/philote_mdo/general/discipline_server.py +++ b/philote_mdo/general/discipline_server.py @@ -27,6 +27,8 @@ # the linked websites, of the information, products, or services contained # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. +import warnings + import grpc import numpy as np @@ -43,67 +45,263 @@ get_variable_shape, read_array_into, ) -from philote_mdo.utils.validation import PhiloteValidationError, validate_shape +from philote_mdo.utils.validation import ( + JobCapacityError, + JobNotFoundError, + JobStateError, + PhiloteJobError, + PhiloteValidationError, + validate_shape, +) +from philote_mdo.general.job import ( + DEFAULT_MAX_JOBS, + DEFAULT_SWEEP_INTERVAL, + DEFAULT_TTL, + JOB_METADATA_KEY, + JobState, + JobStore, +) + + +# gRPC status codes for the job errors raised by JobStore +_JOB_STATUS = ( + (JobNotFoundError, grpc.StatusCode.NOT_FOUND), + (JobCapacityError, grpc.StatusCode.RESOURCE_EXHAUSTED), + (JobStateError, grpc.StatusCode.FAILED_PRECONDITION), +) + + +def _job_status(exc): + """ + Returns the gRPC status code for a job error. + """ + for exc_type, code in _JOB_STATUS: + if isinstance(exc, exc_type): + return code + + return grpc.StatusCode.INTERNAL class DisciplineServer(disc.DisciplineServiceServicer): """ Base class for all server classes. + + The server owns no discipline of its own. It holds a factory and builds one + discipline instance per job, so that the state each client accumulates -- + its options, its variable metadata, and anything the discipline stores on + itself -- stays private to that client. """ - def __init__(self, discipline=None): + def __init__( + self, + discipline=None, + max_jobs=DEFAULT_MAX_JOBS, + ttl=DEFAULT_TTL, + sweep_interval=DEFAULT_SWEEP_INTERVAL, + ): self.verbose = False - # user/developer supplied discipline - self._discipline = discipline + self._max_jobs = max_jobs + self._ttl = ttl + self._sweep_interval = sweep_interval + + # live jobs, created once a factory is attached + self._jobs = None - # discipline stream options - # doubles per message. The cost of a stream is dominated by the number - # of messages in it rather than by their size, so this wants to be as - # large as the message ceiling safely allows: at 100k doubles a chunk - # is about 780 KiB, or a fifth of gRPC's 4 MiB default, which leaves - # room for metadata and for a peer that has lowered the limit. - self._stream_opts = data.StreamOptions(num_double=100000) + if discipline is not None: + self.attach_discipline(discipline) def attach_to_server(self, server): """ Attaches this discipline server class to a gRPC server. """ + self._warn_on_thread_pool(server) disc.add_DisciplineServiceServicer_to_server(self, server) - def attach_discipline(self, impl): + def attach_discipline(self, factory): + """ + Adds a discipline factory to the server. + + Parameters + ---------- + factory : callable + Zero-argument callable returning a fresh discipline. A class works + directly when its ``initialize()`` performs its own configuration; + a discipline configured from the outside needs a closure or a + ``functools.partial``. + """ + self._jobs = JobStore( + factory, + max_jobs=self._max_jobs, + ttl=self._ttl, + sweep_interval=self._sweep_interval, + ) + + def _warn_on_thread_pool(self, server): + """ + Warns when the gRPC thread pool is smaller than the job limit. + + Every in-flight RPC holds a worker for its whole duration, including + the entire bidirectional stream of a compute call. A server that + allows more jobs than the pool has workers queues calls silently + instead of running them, so the limit the operator set is not the one + that binds. + """ + try: + workers = server._state.thread_pool._max_workers + except AttributeError: # pragma: no cover - private gRPC internals + return + + if workers < self._max_jobs: + warnings.warn( + f"gRPC thread pool has {workers} workers but this server " + f"allows {self._max_jobs} concurrent jobs. Every in-flight " + f"RPC holds a worker for its full duration, so jobs beyond " + f"the {workers}th will queue rather than run. Raise " + f"max_workers on the ThreadPoolExecutor, or lower max_jobs.", + RuntimeWarning, + stacklevel=2, + ) + + def _resolve_job(self, context): + """ + Returns the job named by the request metadata. + + Aborts the call rather than returning when the header is missing or + names a job the server does not hold. Call this outside the RPC's + ``try`` block: ``context.abort`` raises, and a surrounding + ``except Exception`` would otherwise swallow it and re-abort with the + wrong status code. + + Parameters + ---------- + context : grpc.ServicerContext + Context of the call in progress. + + Returns + ------- + Job + """ + if self._jobs is None: + context.abort( + grpc.StatusCode.FAILED_PRECONDITION, + "no discipline factory is attached to this server.", + ) + + metadata = dict(context.invocation_metadata() or ()) + job_id = metadata.get(JOB_METADATA_KEY) + + if not job_id: + context.abort( + grpc.StatusCode.FAILED_PRECONDITION, + f"missing '{JOB_METADATA_KEY}' metadata. Call StartJob and " + f"send the returned id on every subsequent call.", + ) + + try: + return self._jobs.get(job_id) + except JobNotFoundError as e: + context.abort(grpc.StatusCode.NOT_FOUND, str(e)) + + def StartJob(self, request, context): """ - Adds a discipline implementation to the server. + RPC that starts a job and returns its handle. """ - self._discipline = impl + if self._jobs is None: + context.abort( + grpc.StatusCode.FAILED_PRECONDITION, + "no discipline factory is attached to this server.", + ) + + try: + job = self._jobs.create() + except PhiloteJobError as e: + context.abort(_job_status(e), str(e)) + except Exception as e: + context.abort( + grpc.StatusCode.INTERNAL, f"StartJob failed: {e}" + ) + + return data.JobHandle(job_id=job.job_id) + + def EndJob(self, request, context): + """ + RPC that ends a job and releases what its discipline holds. + """ + job = self._resolve_job(context) + + try: + self._jobs.close(job.job_id) + return Empty() + except PhiloteJobError as e: + context.abort(_job_status(e), str(e)) + except Exception as e: + context.abort( + grpc.StatusCode.INTERNAL, f"EndJob failed: {e}" + ) + + def KeepAlive(self, request, context): + """ + RPC that defers eviction of an idle job. + """ + # _resolve_job already refreshed the job's timestamp + self._resolve_job(context) + return Empty() + + def _describe(self, context): + """ + Returns a discipline instance for the job-independent RPCs. + + ``GetInfo`` and ``GetAvailableOptions`` report properties of the + discipline class rather than of any run, so they must answer before a + client has a job. Build a throwaway instance for them. + """ + if self._jobs is None: + context.abort( + grpc.StatusCode.FAILED_PRECONDITION, + "no discipline factory is attached to this server.", + ) + + return self._jobs.describe() def GetInfo(self, request, context): """ RPC that sends the discipline information/properties to the client. + + Job-independent: these are properties of the discipline itself. """ + discipline = self._describe(context) + return data.DisciplineProperties( - continuous=self._discipline._is_continuous, - differentiable=self._discipline._is_differentiable, - provides_gradients=self._discipline._provides_gradients, - name=self._discipline._name, - version=self._discipline._version, + continuous=discipline._is_continuous, + differentiable=discipline._is_differentiable, + provides_gradients=discipline._provides_gradients, + name=discipline._name, + version=discipline._version, ) def SetStreamOptions(self, request, context): """ Receives options from the client on how data will be transmitted to and - received from the client. The options are stores locally for use in the - compute routines. + received from the client. The options are stored on the job for use in + the compute routines, since they are a per-client setting. """ - self._stream_opts = request + job = self._resolve_job(context) + job.stream_opts = request return Empty() def GetAvailableOptions(self, request, context): """ RPC that gets the names and types of all available discipline options. + + Job-independent: the option schema comes from ``initialize()`` and is a + property of the discipline class, so a client may call this before it + starts a job. """ + discipline = self._describe(context) + try: - opts_dict = self._discipline.options_list + opts_dict = discipline.options_list opts = data.OptionsList() for name, val in opts_dict.items(): @@ -138,13 +336,23 @@ def GetAvailableOptions(self, request, context): def SetOptions(self, request, context): """ RPC that sets the discipline options. + + Rejected once the job has run ``Setup``: the variable metadata was + built from the previous option values, so accepting new ones would + leave the job describing itself inconsistently. """ + job = self._resolve_job(context) + try: - options = request.options - self._discipline.set_options(options) + with job.lock: + job.require_before_setup("SetOptions") + job.discipline.set_options(request.options) + return Empty() except PhiloteValidationError as e: context.abort(grpc.StatusCode.INVALID_ARGUMENT, str(e)) + except PhiloteJobError as e: + context.abort(_job_status(e), str(e)) except Exception as e: context.abort( grpc.StatusCode.INTERNAL, f"SetOptions failed: {e}" @@ -154,10 +362,16 @@ def Setup(self, request, context): """ RPC that runs the setup function """ + job = self._resolve_job(context) + try: - self._discipline._clear_data() - self._discipline.setup() - self._discipline.setup_partials() + with job.lock: + job.state = JobState.SETUP + job.discipline._clear_data() + job.discipline.setup() + job.discipline.setup_partials() + job.state = JobState.READY + return Empty() except PhiloteValidationError as e: context.abort(grpc.StatusCode.INVALID_ARGUMENT, str(e)) @@ -172,17 +386,21 @@ def GetVariableDefinitions(self, request, context): Both continuous and discrete variable metadata are streamed. """ - for var in self._discipline._var_meta: + job = self._resolve_job(context) + + for var in job.discipline._var_meta: yield var - for var in self._discipline._discrete_var_meta: + for var in job.discipline._discrete_var_meta: yield var def GetPartialDefinitions(self, request, context): """ Transmits partials metadata about the analysis discipline to the client. """ - for jac in self._discipline._partials_meta: + job = self._resolve_job(context) + + for jac in job.discipline._partials_meta: yield jac def SetVariableShapes(self, request_iterator, context): @@ -194,11 +412,13 @@ def SetVariableShapes(self, request_iterator, context): before any compute RPCs for disciplines that contain variables with dynamic shapes. """ + job = self._resolve_job(context) + try: # index by type and name once; searching the list per incoming # message is quadratic in the number of variables index = {} - for var in self._discipline._var_meta: + for var in job.discipline._var_meta: index.setdefault((var.type, var.name), var) for meta in request_iterator: @@ -238,14 +458,21 @@ def SetVariableShapes(self, request_iterator, context): grpc.StatusCode.INTERNAL, f"SetVariableShapes failed: {e}" ) - def preallocate_inputs(self, inputs, flat_inputs, outputs=None, flat_outputs=None): + def preallocate_inputs( + self, job, inputs, flat_inputs, outputs=None, flat_outputs=None + ): """ Preallocates the inputs before receiving data from the client. Note, for implicit disciplines, the function values are considered inputs to evaluate the residuals and the partials of the residuals. + + Parameters + ---------- + job : Job + The job whose discipline supplies the variable metadata. """ - for var in self._discipline._var_meta: + for var in job.discipline._var_meta: # validate that dynamic-shape variables have been resolved if var.dynamic_shape and len(var.shape) == 0: raise PhiloteValidationError( @@ -266,21 +493,26 @@ def preallocate_inputs(self, inputs, flat_inputs, outputs=None, flat_outputs=Non outputs[var.name] = np.zeros(var.shape) flat_outputs[var.name] = get_flattened_view(outputs[var.name]) - def preallocate_partials(self): + def preallocate_partials(self, job): """ Preallocates the partials. Note: there are edge cases for this function, where either f or x, or both are scalar. In those cases the shapes of the partials must be treated differently. + + Parameters + ---------- + job : Job + The job whose discipline supplies the partials metadata. """ jac = PairDict() # index the metadata by (type, name) once; scanning it per partial is # quadratic in the number of variables - shapes = build_shape_index(self._discipline._var_meta) + shapes = build_shape_index(job.discipline._var_meta) - for pair in self._discipline._partials_meta: + for pair in job.discipline._partials_meta: shape = get_partials_shape( get_function_shape(shapes, pair.name, "preallocate_partials"), get_variable_shape(shapes, pair.subname, "preallocate_partials"), diff --git a/philote_mdo/general/explicit_client.py b/philote_mdo/general/explicit_client.py index b42f8e7..2b5f1ae 100644 --- a/philote_mdo/general/explicit_client.py +++ b/philote_mdo/general/explicit_client.py @@ -29,7 +29,8 @@ # control over the information you may find at these locations. import grpc from philote_mdo.general.discipline_client import DisciplineClient -from philote_mdo.utils.validation import PhiloteServerError, validate_is_dict +from philote_mdo.general.discipline_client import raise_for_rpc_error +from philote_mdo.utils.validation import validate_is_dict import philote_mdo.generated.disciplines_pb2_grpc as disc @@ -40,7 +41,7 @@ class ExplicitClient(DisciplineClient): def __init__(self, channel): super().__init__(channel) - self._expl_stub = disc.ExplicitServiceStub(channel) + self._expl_stub = disc.ExplicitServiceStub(self._channel) def run_compute(self, inputs, discrete_inputs=None): """ @@ -60,6 +61,7 @@ def run_compute(self, inputs, discrete_inputs=None): Continuous outputs, or (continuous outputs, discrete outputs) when the server returns discrete output data. """ + self._ensure_job() validate_is_dict(inputs, "run_compute (inputs)") try: messages = self._assemble_input_messages( @@ -68,9 +70,7 @@ def run_compute(self, inputs, discrete_inputs=None): responses = self._expl_stub.ComputeFunction(iter(messages)) return self._recover_outputs(responses) except grpc.RpcError as e: - raise PhiloteServerError( - f"Server error during run_compute: {e.details()}" - ) from e + raise_for_rpc_error(e, "run_compute") def run_compute_partials(self, inputs, discrete_inputs=None): """ @@ -84,6 +84,7 @@ def run_compute_partials(self, inputs, discrete_inputs=None): discrete_inputs : dict, optional Discrete input values. """ + self._ensure_job() validate_is_dict(inputs, "run_compute_partials (inputs)") try: messages = self._assemble_input_messages( @@ -94,6 +95,4 @@ def run_compute_partials(self, inputs, discrete_inputs=None): return partials except grpc.RpcError as e: - raise PhiloteServerError( - f"Server error during run_compute_partials: {e.details()}" - ) from e + raise_for_rpc_error(e, "run_compute_partials") diff --git a/philote_mdo/general/explicit_server.py b/philote_mdo/general/explicit_server.py index 7129e13..e5bd087 100644 --- a/philote_mdo/general/explicit_server.py +++ b/philote_mdo/general/explicit_server.py @@ -43,8 +43,8 @@ class ExplicitServer(DisciplineServer, disc.ExplicitServiceServicer): Base class for remote explicit components. """ - def __init__(self, discipline=None): - super().__init__(discipline=discipline) + def __init__(self, discipline=None, **kwargs): + super().__init__(discipline=discipline, **kwargs) def attach_to_server(self, server): """ @@ -57,6 +57,12 @@ def ComputeFunction(self, request_iterator, context): """ Computes the function evaluation and sends the result to the client. """ + job = self._resolve_job(context) + + # serialise calls within this job. Separate jobs never contend here, + # which is what lets two clients evaluate at the same time. + job.lock.acquire() + try: inputs = {} flat_inputs = {} @@ -64,22 +70,22 @@ def ComputeFunction(self, request_iterator, context): discrete_inputs = {} discrete_outputs = {} - self.preallocate_inputs(inputs, flat_inputs) + self.preallocate_inputs(job, inputs, flat_inputs) discrete_inputs, _ = self.process_inputs( request_iterator, flat_inputs, discrete_inputs=discrete_inputs ) # Call compute with discrete data when discrete variables are present - if discrete_inputs or self._discipline._discrete_var_meta: - self._discipline.compute( + if discrete_inputs or job.discipline._discrete_var_meta: + job.discipline.compute( inputs, outputs, discrete_inputs, discrete_outputs ) else: - self._discipline.compute(inputs, outputs) + job.discipline.compute(inputs, outputs) # Stream continuous outputs for output_name, value in outputs.items(): - for b, e in get_chunk_indices(value.size, self._stream_opts.num_double): + for b, e in get_chunk_indices(value.size, job.stream_opts.num_double): message = data.VariableMessage( continuous=data.Array( name=output_name, @@ -107,29 +113,37 @@ def ComputeFunction(self, request_iterator, context): context.abort( grpc.StatusCode.INTERNAL, f"ComputeFunction failed: {e}" ) + finally: + job.lock.release() def ComputeGradient(self, request_iterator, context): """ Computes the gradient evaluation and sends the result to the client. """ + job = self._resolve_job(context) + + # serialise calls within this job. Separate jobs never contend here, + # which is what lets two clients evaluate at the same time. + job.lock.acquire() + try: inputs = {} flat_inputs = {} discrete_inputs = {} - self.preallocate_inputs(inputs, flat_inputs) - jac = self.preallocate_partials() + self.preallocate_inputs(job, inputs, flat_inputs) + jac = self.preallocate_partials(job) discrete_inputs, _ = self.process_inputs( request_iterator, flat_inputs, discrete_inputs=discrete_inputs ) - if discrete_inputs or self._discipline._discrete_var_meta: - self._discipline.compute_partials(inputs, jac, discrete_inputs) + if discrete_inputs or job.discipline._discrete_var_meta: + job.discipline.compute_partials(inputs, jac, discrete_inputs) else: - self._discipline.compute_partials(inputs, jac) + job.discipline.compute_partials(inputs, jac) for jac, value in jac.items(): - for b, e in get_chunk_indices(value.size, self._stream_opts.num_double): + for b, e in get_chunk_indices(value.size, job.stream_opts.num_double): message = data.VariableMessage( continuous=data.Array( name=jac[0], @@ -148,3 +162,5 @@ def ComputeGradient(self, request_iterator, context): context.abort( grpc.StatusCode.INTERNAL, f"ComputeGradient failed: {e}" ) + finally: + job.lock.release() diff --git a/philote_mdo/general/implicit_client.py b/philote_mdo/general/implicit_client.py index 9ef8156..af77a3d 100644 --- a/philote_mdo/general/implicit_client.py +++ b/philote_mdo/general/implicit_client.py @@ -29,7 +29,8 @@ # control over the information you may find at these locations. import grpc from philote_mdo.general.discipline_client import DisciplineClient -from philote_mdo.utils.validation import PhiloteServerError, validate_is_dict +from philote_mdo.general.discipline_client import raise_for_rpc_error +from philote_mdo.utils.validation import validate_is_dict import philote_mdo.generated.data_pb2 as data import philote_mdo.generated.disciplines_pb2_grpc as disc @@ -118,7 +119,7 @@ def __init__(self, channel): - Variable metadata and options are automatically discovered from server """ super().__init__(channel=channel) - self._impl_stub = disc.ImplicitServiceStub(channel) + self._impl_stub = disc.ImplicitServiceStub(self._channel) def run_compute_residuals( self, inputs, outputs, discrete_inputs=None, discrete_outputs=None @@ -169,6 +170,7 @@ def run_compute_residuals( - Large arrays are automatically streamed for efficiency - This is typically used for residual evaluation during Newton iterations """ + self._ensure_job() validate_is_dict(inputs, "run_compute_residuals (inputs)") validate_is_dict(outputs, "run_compute_residuals (outputs)") try: @@ -184,9 +186,7 @@ def run_compute_residuals( return residuals except grpc.RpcError as e: - raise PhiloteServerError( - f"Server error during run_compute_residuals: {e.details()}" - ) from e + raise_for_rpc_error(e, "run_compute_residuals") def run_solve_residuals(self, inputs, discrete_inputs=None): """ @@ -239,6 +239,7 @@ def run_solve_residuals(self, inputs, discrete_inputs=None): - May raise exceptions for ill-conditioned or non-convergent problems - Solution quality depends on the server's implementation and input conditioning """ + self._ensure_job() validate_is_dict(inputs, "run_solve_residuals (inputs)") try: # Assemble input messages and call server @@ -249,9 +250,7 @@ def run_solve_residuals(self, inputs, discrete_inputs=None): outputs = self._recover_outputs(responses) return outputs except grpc.RpcError as e: - raise PhiloteServerError( - f"Server error during run_solve_residuals: {e.details()}" - ) from e + raise_for_rpc_error(e, "run_solve_residuals") def run_residual_gradients( self, inputs, outputs, discrete_inputs=None, discrete_outputs=None @@ -313,6 +312,7 @@ def run_residual_gradients( - Used by optimization algorithms and sensitivity analysis tools - For large problems, consider matrix-free methods if available """ + self._ensure_job() validate_is_dict(inputs, "run_residual_gradients (inputs)") validate_is_dict(outputs, "run_residual_gradients (outputs)") try: @@ -327,6 +327,4 @@ def run_residual_gradients( partials = self._recover_partials(responses) return partials except grpc.RpcError as e: - raise PhiloteServerError( - f"Server error during run_residual_gradients: {e.details()}" - ) from e + raise_for_rpc_error(e, "run_residual_gradients") diff --git a/philote_mdo/general/implicit_server.py b/philote_mdo/general/implicit_server.py index 21c3f64..835bb03 100644 --- a/philote_mdo/general/implicit_server.py +++ b/philote_mdo/general/implicit_server.py @@ -64,12 +64,12 @@ class ImplicitServer(pmdo.DisciplineServer, disc.ImplicitServiceServicer): >>> import grpc >>> import philote_mdo.general as pmdo >>> - >>> # Create your implicit discipline - >>> discipline = MyImplicitDiscipline() - >>> - >>> # Create and configure server - >>> server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - >>> impl_server = pmdo.ImplicitServer(discipline=discipline) + >>> # Create and configure server. One discipline instance is built + >>> # per job, so concurrent clients do not interfere. + >>> server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) + >>> impl_server = pmdo.ImplicitServer( + ... discipline=MyImplicitDiscipline + ... ) >>> impl_server.attach_to_server(server) >>> >>> # Start server @@ -79,7 +79,7 @@ class ImplicitServer(pmdo.DisciplineServer, disc.ImplicitServiceServicer): >>> server.wait_for_termination() Attributes: - _discipline (ImplicitDiscipline): The underlying implicit discipline being served + _jobs (JobStore): The live jobs, each owning its own discipline instance Notes: - Inherits from both DisciplineServer and gRPC ImplicitServiceServicer @@ -88,22 +88,26 @@ class ImplicitServer(pmdo.DisciplineServer, disc.ImplicitServiceServicer): - The underlying discipline must implement all required implicit methods """ - def __init__(self, discipline=None): + def __init__(self, discipline=None, **kwargs): """ Initialize the implicit discipline server. Parameters ---------- - discipline : ImplicitDiscipline, optional - The implicit discipline to serve. Must implement compute_residuals, - solve_residuals, and residual_partials methods. + discipline : callable, optional + Zero-argument callable returning a fresh implicit discipline, which + must implement compute_residuals, solve_residuals and + residual_partials. One instance is built per job, so the state a + client accumulates stays private to it. + **kwargs + Job limits passed to ``DisciplineServer``: ``max_jobs``, ``ttl`` + and ``sweep_interval``. Examples -------- - >>> discipline = MyImplicitDiscipline() - >>> server = ImplicitServer(discipline=discipline) + >>> server = ImplicitServer(discipline=MyImplicitDiscipline) """ - super().__init__(discipline=discipline) + super().__init__(discipline=discipline, **kwargs) def attach_to_server(self, server): """ @@ -120,7 +124,7 @@ def attach_to_server(self, server): Examples -------- >>> server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - >>> impl_server = ImplicitServer(discipline=my_discipline) + >>> impl_server = ImplicitServer(discipline=MyImplicitDiscipline) >>> impl_server.attach_to_server(server) >>> server.add_insecure_port('[::]:50051') >>> server.start() @@ -156,6 +160,12 @@ def ComputeResiduals(self, request_iterator, context): - Streams results back in chunks for efficiency - This method is called automatically by the gRPC framework """ + job = self._resolve_job(context) + + # serialise calls within this job. Separate jobs never contend here, + # which is what lets two clients evaluate at the same time. + job.lock.acquire() + try: # inputs and outputs inputs = {} @@ -166,7 +176,7 @@ def ComputeResiduals(self, request_iterator, context): discrete_inputs = {} discrete_outputs = {} - self.preallocate_inputs(inputs, flat_inputs, outputs, flat_outputs) + self.preallocate_inputs(job, inputs, flat_inputs, outputs, flat_outputs) discrete_inputs, discrete_outputs = self.process_inputs( request_iterator, flat_inputs, @@ -176,15 +186,15 @@ def ComputeResiduals(self, request_iterator, context): ) # Call the user-defined compute_residuals function - if discrete_inputs or self._discipline._discrete_var_meta: - self._discipline.compute_residuals( + if discrete_inputs or job.discipline._discrete_var_meta: + job.discipline.compute_residuals( inputs, outputs, residuals, discrete_inputs, discrete_outputs ) else: - self._discipline.compute_residuals(inputs, outputs, residuals) + job.discipline.compute_residuals(inputs, outputs, residuals) for res_name, value in residuals.items(): - for b, e in get_chunk_indices(value.size, self._stream_opts.num_double): + for b, e in get_chunk_indices(value.size, job.stream_opts.num_double): message = data.VariableMessage( continuous=data.Array( name=res_name, @@ -202,6 +212,8 @@ def ComputeResiduals(self, request_iterator, context): context.abort( grpc.StatusCode.INTERNAL, f"ComputeResiduals failed: {e}" ) + finally: + job.lock.release() def SolveResiduals(self, request_iterator, context): """ @@ -230,6 +242,12 @@ def SolveResiduals(self, request_iterator, context): - Outputs are streamed back in chunks for large arrays - This method is called automatically by the gRPC framework """ + job = self._resolve_job(context) + + # serialise calls within this job. Separate jobs never contend here, + # which is what lets two clients evaluate at the same time. + job.lock.acquire() + try: # inputs and outputs inputs = {} @@ -238,7 +256,7 @@ def SolveResiduals(self, request_iterator, context): flat_outputs = {} discrete_inputs = {} - self.preallocate_inputs(inputs, flat_inputs, outputs, flat_outputs) + self.preallocate_inputs(job, inputs, flat_inputs, outputs, flat_outputs) discrete_inputs, _ = self.process_inputs( request_iterator, flat_inputs, @@ -247,13 +265,13 @@ def SolveResiduals(self, request_iterator, context): ) # Call the user-defined solve function - if discrete_inputs or self._discipline._discrete_var_meta: - self._discipline.solve_residuals(inputs, outputs, discrete_inputs) + if discrete_inputs or job.discipline._discrete_var_meta: + job.discipline.solve_residuals(inputs, outputs, discrete_inputs) else: - self._discipline.solve_residuals(inputs, outputs) + job.discipline.solve_residuals(inputs, outputs) for output_name, value in outputs.items(): - for b, e in get_chunk_indices(value.size, self._stream_opts.num_double): + for b, e in get_chunk_indices(value.size, job.stream_opts.num_double): message = data.VariableMessage( continuous=data.Array( name=output_name, @@ -271,6 +289,8 @@ def SolveResiduals(self, request_iterator, context): context.abort( grpc.StatusCode.INTERNAL, f"SolveResiduals failed: {e}" ) + finally: + job.lock.release() def ComputeResidualGradients(self, request_iterator, context): """ @@ -299,6 +319,12 @@ def ComputeResidualGradients(self, request_iterator, context): - Used for gradient-based optimization and sensitivity analysis - This method is called automatically by the gRPC framework """ + job = self._resolve_job(context) + + # serialise calls within this job. Separate jobs never contend here, + # which is what lets two clients evaluate at the same time. + job.lock.acquire() + try: # inputs and outputs inputs = {} @@ -308,8 +334,8 @@ def ComputeResidualGradients(self, request_iterator, context): discrete_inputs = {} discrete_outputs = {} - self.preallocate_inputs(inputs, flat_inputs, outputs, flat_outputs) - jac = self.preallocate_partials() + self.preallocate_inputs(job, inputs, flat_inputs, outputs, flat_outputs) + jac = self.preallocate_partials(job) discrete_inputs, discrete_outputs = self.process_inputs( request_iterator, flat_inputs, @@ -319,15 +345,15 @@ def ComputeResidualGradients(self, request_iterator, context): ) # Call the user-defined residual partials function - if discrete_inputs or self._discipline._discrete_var_meta: - self._discipline.residual_partials( + if discrete_inputs or job.discipline._discrete_var_meta: + job.discipline.residual_partials( inputs, outputs, jac, discrete_inputs, discrete_outputs ) else: - self._discipline.residual_partials(inputs, outputs, jac) + job.discipline.residual_partials(inputs, outputs, jac) for jac, value in jac.items(): - for b, e in get_chunk_indices(value.size, self._stream_opts.num_double): + for b, e in get_chunk_indices(value.size, job.stream_opts.num_double): message = data.VariableMessage( continuous=data.Array( name=jac[0], @@ -347,6 +373,8 @@ def ComputeResidualGradients(self, request_iterator, context): grpc.StatusCode.INTERNAL, f"ComputeResidualGradients failed: {e}", ) + finally: + job.lock.release() # def MatrixFreeGradients(self, request_iterator, context): # """ diff --git a/philote_mdo/general/job.py b/philote_mdo/general/job.py new file mode 100644 index 0000000..490c012 --- /dev/null +++ b/philote_mdo/general/job.py @@ -0,0 +1,406 @@ +# Philote-Python +# +# Copyright 2022-2025 Christopher A. Lupp +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# +# This work has been cleared for public release, distribution unlimited, case +# number: AFRL-2023-5713. +# +# The views expressed are those of the authors and do not reflect the +# official guidance or position of the United States Government, the +# Department of Defense or of the United States Air Force. +# +# Statement from DoD: The Appearance of external hyperlinks does not +# constitute endorsement by the United States Department of Defense (DoD) of +# the linked websites, of the information, products, or services contained +# therein. The DoD does not exercise any editorial, security, or other +# control over the information you may find at these locations. +""" +Server-side job sessions. + +A Philote server used to hold one discipline instance and share it across +every client that connected. Because the protocol has each client run a setup +sequence that mutates that instance -- ``SetOptions``, ``Setup``, and the +in-place shape edits of ``SetVariableShapes`` -- two concurrent clients +corrupted one another. + +A :class:`Job` is a session that owns its own discipline instance, so the state +each client builds is private to it. Clients get a job id from ``StartJob`` and +present it in the ``philote-job-id`` metadata header on every later call. + +Jobs are session-shaped rather than request-shaped: there are few of them, each +lives across many evaluations, and each can hold something substantial such as +an initialised solver or an ``om.Problem``. That is why they are capped and +swept rather than allowed to accumulate. +""" +import threading +import time +import uuid + +import philote_mdo.generated.data_pb2 as data +from philote_mdo.utils.validation import ( + JobCapacityError, + JobNotFoundError, + JobStateError, +) + + +# metadata header carrying the job id on every job-scoped RPC +JOB_METADATA_KEY = "philote-job-id" + +# doubles per message, matching the server default. Held per job because +# SetStreamOptions is a per-client setting. +DEFAULT_NUM_DOUBLE = 100000 + +# a job that has gone unused for this many seconds is evicted +DEFAULT_TTL = 3600.0 + +# how often the sweeper thread looks for expired jobs +DEFAULT_SWEEP_INTERVAL = 60.0 + +# concurrent jobs allowed per server. Deliberately modest: each job holds a +# discipline instance, and the gRPC thread pool caps useful concurrency anyway. +DEFAULT_MAX_JOBS = 8 + + +class JobState: + """The stages a job moves through, in order.""" + + NEW = "new" + SETUP = "setup" + READY = "ready" + CLOSED = "closed" + + +class Job: + """One client session, owning one discipline instance. + + Parameters + ---------- + job_id : str + Server-assigned identifier, opaque to the client. + discipline : Discipline or None + The discipline instance this job owns. ``None`` only while + :meth:`JobStore.create` is still building it. + """ + + def __init__(self, job_id, discipline=None): + self.job_id = job_id + self.discipline = discipline + + # stream options are per client, so they live here rather than on the + # server that all clients share + self.stream_opts = data.StreamOptions(num_double=DEFAULT_NUM_DOUBLE) + + self.state = JobState.NEW + + # serialises calls within this job. Separate jobs never contend on it, + # so two clients can be inside compute() at the same time. + self.lock = threading.Lock() + + self.last_used = time.monotonic() + + def touch(self): + """Marks the job as still in use, deferring eviction.""" + self.last_used = time.monotonic() + + def require_before_setup(self, rpc_name): + """Rejects a call that must precede ``Setup`` but arrived after it. + + Parameters + ---------- + rpc_name : str + Name of the RPC, used in the error message. + + Raises + ------ + JobStateError + If ``Setup`` has already run for this job. + """ + if self.state in (JobState.SETUP, JobState.READY): + raise JobStateError( + f"{rpc_name}: job '{self.job_id}' has already run Setup. " + f"The variable metadata was built from the previous values, " + f"so start a new job instead." + ) + + def __repr__(self): + return f"" + + +class JobStore: + """Holds the live jobs for one server and enforces their limits. + + Parameters + ---------- + discipline_factory : callable + Zero-argument callable returning a fresh ``Discipline``. A class works + directly when its ``initialize()`` does its own configuration; a + discipline configured from outside needs a closure or + ``functools.partial``. + max_jobs : int, optional + Concurrent jobs allowed before ``StartJob`` is refused. + ttl : float or None, optional + Seconds a job may sit unused before eviction. ``None`` disables both + expiry and the sweeper thread. + sweep_interval : float, optional + Seconds between sweeps. Ignored when ``ttl`` is ``None``. + """ + + def __init__( + self, + discipline_factory, + max_jobs=DEFAULT_MAX_JOBS, + ttl=DEFAULT_TTL, + sweep_interval=DEFAULT_SWEEP_INTERVAL, + ): + if not callable(discipline_factory): + raise TypeError( + f"discipline must be a zero-argument callable returning a " + f"Discipline, got an instance of " + f"{type(discipline_factory).__name__}. Pass the class rather " + f"than an instance of it -- " + f"ExplicitServer(discipline={type(discipline_factory).__name__})" + f", not " + f"ExplicitServer(discipline={type(discipline_factory).__name__}())" + f". The server builds one discipline per job, so it needs " + f"something it can call." + ) + + self._factory = discipline_factory + self._max_jobs = max_jobs + self._ttl = ttl + self._jobs = {} + self._lock = threading.Lock() + + # lazily built instance answering the job-independent RPCs + self._prototype = None + + self._stop = threading.Event() + self._sweeper = None + + if ttl is not None: + self._sweeper = threading.Thread( + target=self._sweep_loop, + args=(sweep_interval,), + name="philote-job-sweeper", + daemon=True, + ) + self._sweeper.start() + + @property + def max_jobs(self): + return self._max_jobs + + def describe(self): + """Returns an instance for the RPCs that do not belong to a job. + + ``GetInfo`` and ``GetAvailableOptions`` report properties of the + discipline class -- its name, version, and the option schema built by + ``initialize()`` -- so a client must be able to call them before it has + a job. One instance is built on first use and reused, and it is never + handed to a job, so nothing a client does can reach it. + + Returns + ------- + Discipline + """ + with self._lock: + if self._prototype is None: + self._prototype = self._factory() + + return self._prototype + + def create(self): + """Starts a job and builds its discipline. + + Returns + ------- + Job + The new job, with its discipline built and back-linked. + + Raises + ------ + JobCapacityError + If the server already holds ``max_jobs`` jobs. + """ + # reserve the slot under the store lock, then build the discipline + # outside it. Construction can be slow -- loading a mesh, building an + # om.Problem -- and holding the store lock through it would stall every + # other job's calls. + expired = [] + + try: + with self._lock: + # reclaim expired slots first, so a dead client cannot make + # StartJob fail while the sweeper has yet to run + expired = self._expired_locked() + + for old in expired: + self._jobs.pop(old.job_id, None) + + if len(self._jobs) >= self._max_jobs: + raise JobCapacityError( + f"StartJob: server already holds its maximum of " + f"{self._max_jobs} jobs. End a job before starting " + f"another, or raise max_jobs." + ) + + job_id = uuid.uuid4().hex + job = Job(job_id) + self._jobs[job_id] = job + finally: + # outside the lock, and on the capacity path too: whatever this + # call evicted still has to release what it held + for old in expired: + self._teardown(old) + + try: + discipline = self._factory() + except Exception: + with self._lock: + self._jobs.pop(job_id, None) + raise + + job.discipline = discipline + + # let the discipline reach its own job, for the job id and, once the + # file API lands, its scratch directory + discipline.job = job + + return job + + def get(self, job_id): + """Looks up a live job and marks it as used. + + Parameters + ---------- + job_id : str + The id from the ``philote-job-id`` header. + + Returns + ------- + Job + + Raises + ------ + JobNotFoundError + If no such job exists, or it has been closed or evicted. + """ + with self._lock: + job = self._jobs.get(job_id) + + if job is None: + raise JobNotFoundError( + f"job '{job_id}' is unknown to this server. It may have " + f"been ended, expired, or belonged to a server that has " + f"since restarted. Start a new job; any state it held is " + f"gone." + ) + + job.touch() + + return job + + def close(self, job_id): + """Ends a job and releases whatever its discipline holds. + + Parameters + ---------- + job_id : str + The job to end. + + Raises + ------ + JobNotFoundError + If no such job exists. + """ + with self._lock: + job = self._jobs.pop(job_id, None) + + if job is None: + raise JobNotFoundError( + f"EndJob: job '{job_id}' is unknown to this server." + ) + + self._teardown(job) + + def sweep(self): + """Evicts every job that has outlived the TTL. + + Called by the sweeper thread, and directly by tests. + + Returns + ------- + list of str + The ids that were evicted. + """ + with self._lock: + expired = self._expired_locked() + + for job in expired: + self._jobs.pop(job.job_id, None) + + for job in expired: + self._teardown(job) + + return [job.job_id for job in expired] + + def close_all(self): + """Ends every job and stops the sweeper. Used on server shutdown.""" + self._stop.set() + + with self._lock: + jobs = list(self._jobs.values()) + self._jobs.clear() + + for job in jobs: + self._teardown(job) + + def __len__(self): + with self._lock: + return len(self._jobs) + + def _expired_locked(self): + """Returns the expired jobs. The store lock must be held.""" + if self._ttl is None: + return [] + + cutoff = time.monotonic() - self._ttl + + return [job for job in self._jobs.values() if job.last_used < cutoff] + + def _teardown(self, job): + """Lets the discipline release its resources, then marks it closed.""" + job.state = JobState.CLOSED + + discipline = job.discipline + + if discipline is None: + return + + teardown = getattr(discipline, "teardown_job", None) + + if callable(teardown): + teardown() + + job.discipline = None + + def _sweep_loop(self, interval): + while not self._stop.wait(interval): + try: + self.sweep() + except Exception: # pragma: no cover - a sweep must never kill + pass # the thread; the next tick tries again diff --git a/philote_mdo/generated/data_pb2.py b/philote_mdo/generated/data_pb2.py index 668e1ab..d123bfe 100644 --- a/philote_mdo/generated/data_pb2.py +++ b/philote_mdo/generated/data_pb2.py @@ -7,32 +7,34 @@ _runtime_version.ValidateProtobufRuntimeVersion(_runtime_version.Domain.PUBLIC, 5, 27, 2, '', 'data.proto') _sym_db = _symbol_database.Default() from google.protobuf import struct_pb2 as google_dot_protobuf_dot_struct__pb2 -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\ndata.proto\x12\x07philote\x1a\x1cgoogle/protobuf/struct.proto"}\n\x14DisciplineProperties\x12\x12\n\ncontinuous\x18\x01 \x01(\x08\x12\x16\n\x0edifferentiable\x18\x02 \x01(\x08\x12\x1a\n\x12provides_gradients\x18\x03 \x01(\x08\x12\x0c\n\x04name\x18\x04 \x01(\t\x12\x0f\n\x07version\x18\x05 \x01(\t"#\n\rStreamOptions\x12\x12\n\nnum_double\x18\x01 \x01(\x03"?\n\x0bOptionsList\x12\x0f\n\x07options\x18\x01 \x03(\t\x12\x1f\n\x04type\x18\x02 \x03(\x0e2\x11.philote.DataType"=\n\x11DisciplineOptions\x12(\n\x07options\x18\x01 \x01(\x0b2\x17.google.protobuf.Struct"z\n\x10VariableMetaData\x12#\n\x04type\x18\x01 \x01(\x0e2\x15.philote.VariableType\x12\x0c\n\x04name\x18\x03 \x01(\t\x12\r\n\x05shape\x18\x04 \x03(\x03\x12\r\n\x05units\x18\x05 \x01(\t\x12\x15\n\rdynamic_shape\x18\x06 \x01(\x08"@\n\x10PartialsMetaData\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0f\n\x07subname\x18\x02 \x01(\t\x12\r\n\x05shape\x18\x03 \x03(\x03"u\n\x05Array\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0f\n\x07subname\x18\x02 \x01(\t\x12\r\n\x05start\x18\x03 \x01(\x03\x12\x0b\n\x03end\x18\x04 \x01(\x03\x12#\n\x04type\x18\x05 \x01(\x0e2\x15.philote.VariableType\x12\x0c\n\x04data\x18\x06 \x03(\x01"l\n\x10DiscreteVariable\x12\x0c\n\x04name\x18\x01 \x01(\t\x12#\n\x04type\x18\x02 \x01(\x0e2\x15.philote.VariableType\x12%\n\x05value\x18\x03 \x01(\x0b2\x16.google.protobuf.Value"q\n\x0fVariableMessage\x12$\n\ncontinuous\x18\x01 \x01(\x0b2\x0e.philote.ArrayH\x00\x12-\n\x08discrete\x18\x02 \x01(\x0b2\x19.philote.DiscreteVariableH\x00B\t\n\x07payload*F\n\x08DataType\x12\t\n\x05kBool\x10\x00\x12\x08\n\x04kInt\x10\x01\x12\x0b\n\x07kDouble\x10\x02\x12\x0b\n\x07kString\x10\x03\x12\x0b\n\x07kStruct\x10\x04*m\n\x0cVariableType\x12\n\n\x06kInput\x10\x00\x12\x12\n\x0ekDiscreteInput\x10\x01\x12\r\n\tkResidual\x10\x02\x12\x0b\n\x07kOutput\x10\x03\x12\x13\n\x0fkDiscreteOutput\x10\x04\x12\x0c\n\x08kPartial\x10\x05B\x11\n\x0forg.philote.mdob\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\ndata.proto\x12\x07philote\x1a\x1cgoogle/protobuf/struct.proto"}\n\x14DisciplineProperties\x12\x12\n\ncontinuous\x18\x01 \x01(\x08\x12\x16\n\x0edifferentiable\x18\x02 \x01(\x08\x12\x1a\n\x12provides_gradients\x18\x03 \x01(\x08\x12\x0c\n\x04name\x18\x04 \x01(\t\x12\x0f\n\x07version\x18\x05 \x01(\t"\x1b\n\tJobHandle\x12\x0e\n\x06job_id\x18\x01 \x01(\t"#\n\rStreamOptions\x12\x12\n\nnum_double\x18\x01 \x01(\x03"?\n\x0bOptionsList\x12\x0f\n\x07options\x18\x01 \x03(\t\x12\x1f\n\x04type\x18\x02 \x03(\x0e2\x11.philote.DataType"=\n\x11DisciplineOptions\x12(\n\x07options\x18\x01 \x01(\x0b2\x17.google.protobuf.Struct"z\n\x10VariableMetaData\x12#\n\x04type\x18\x01 \x01(\x0e2\x15.philote.VariableType\x12\x0c\n\x04name\x18\x03 \x01(\t\x12\r\n\x05shape\x18\x04 \x03(\x03\x12\r\n\x05units\x18\x05 \x01(\t\x12\x15\n\rdynamic_shape\x18\x06 \x01(\x08"@\n\x10PartialsMetaData\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0f\n\x07subname\x18\x02 \x01(\t\x12\r\n\x05shape\x18\x03 \x03(\x03"u\n\x05Array\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0f\n\x07subname\x18\x02 \x01(\t\x12\r\n\x05start\x18\x03 \x01(\x03\x12\x0b\n\x03end\x18\x04 \x01(\x03\x12#\n\x04type\x18\x05 \x01(\x0e2\x15.philote.VariableType\x12\x0c\n\x04data\x18\x06 \x03(\x01"l\n\x10DiscreteVariable\x12\x0c\n\x04name\x18\x01 \x01(\t\x12#\n\x04type\x18\x02 \x01(\x0e2\x15.philote.VariableType\x12%\n\x05value\x18\x03 \x01(\x0b2\x16.google.protobuf.Value"q\n\x0fVariableMessage\x12$\n\ncontinuous\x18\x01 \x01(\x0b2\x0e.philote.ArrayH\x00\x12-\n\x08discrete\x18\x02 \x01(\x0b2\x19.philote.DiscreteVariableH\x00B\t\n\x07payload*F\n\x08DataType\x12\t\n\x05kBool\x10\x00\x12\x08\n\x04kInt\x10\x01\x12\x0b\n\x07kDouble\x10\x02\x12\x0b\n\x07kString\x10\x03\x12\x0b\n\x07kStruct\x10\x04*m\n\x0cVariableType\x12\n\n\x06kInput\x10\x00\x12\x12\n\x0ekDiscreteInput\x10\x01\x12\r\n\tkResidual\x10\x02\x12\x0b\n\x07kOutput\x10\x03\x12\x13\n\x0fkDiscreteOutput\x10\x04\x12\x0c\n\x08kPartial\x10\x05B\x11\n\x0forg.philote.mdob\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) _builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'data_pb2', _globals) if not _descriptor._USE_C_DESCRIPTORS: _globals['DESCRIPTOR']._loaded_options = None _globals['DESCRIPTOR']._serialized_options = b'\n\x0forg.philote.mdo' - _globals['_DATATYPE']._serialized_start = 879 - _globals['_DATATYPE']._serialized_end = 949 - _globals['_VARIABLETYPE']._serialized_start = 951 - _globals['_VARIABLETYPE']._serialized_end = 1060 + _globals['_DATATYPE']._serialized_start = 908 + _globals['_DATATYPE']._serialized_end = 978 + _globals['_VARIABLETYPE']._serialized_start = 980 + _globals['_VARIABLETYPE']._serialized_end = 1089 _globals['_DISCIPLINEPROPERTIES']._serialized_start = 53 _globals['_DISCIPLINEPROPERTIES']._serialized_end = 178 - _globals['_STREAMOPTIONS']._serialized_start = 180 - _globals['_STREAMOPTIONS']._serialized_end = 215 - _globals['_OPTIONSLIST']._serialized_start = 217 - _globals['_OPTIONSLIST']._serialized_end = 280 - _globals['_DISCIPLINEOPTIONS']._serialized_start = 282 - _globals['_DISCIPLINEOPTIONS']._serialized_end = 343 - _globals['_VARIABLEMETADATA']._serialized_start = 345 - _globals['_VARIABLEMETADATA']._serialized_end = 467 - _globals['_PARTIALSMETADATA']._serialized_start = 469 - _globals['_PARTIALSMETADATA']._serialized_end = 533 - _globals['_ARRAY']._serialized_start = 535 - _globals['_ARRAY']._serialized_end = 652 - _globals['_DISCRETEVARIABLE']._serialized_start = 654 - _globals['_DISCRETEVARIABLE']._serialized_end = 762 - _globals['_VARIABLEMESSAGE']._serialized_start = 764 - _globals['_VARIABLEMESSAGE']._serialized_end = 877 \ No newline at end of file + _globals['_JOBHANDLE']._serialized_start = 180 + _globals['_JOBHANDLE']._serialized_end = 207 + _globals['_STREAMOPTIONS']._serialized_start = 209 + _globals['_STREAMOPTIONS']._serialized_end = 244 + _globals['_OPTIONSLIST']._serialized_start = 246 + _globals['_OPTIONSLIST']._serialized_end = 309 + _globals['_DISCIPLINEOPTIONS']._serialized_start = 311 + _globals['_DISCIPLINEOPTIONS']._serialized_end = 372 + _globals['_VARIABLEMETADATA']._serialized_start = 374 + _globals['_VARIABLEMETADATA']._serialized_end = 496 + _globals['_PARTIALSMETADATA']._serialized_start = 498 + _globals['_PARTIALSMETADATA']._serialized_end = 562 + _globals['_ARRAY']._serialized_start = 564 + _globals['_ARRAY']._serialized_end = 681 + _globals['_DISCRETEVARIABLE']._serialized_start = 683 + _globals['_DISCRETEVARIABLE']._serialized_end = 791 + _globals['_VARIABLEMESSAGE']._serialized_start = 793 + _globals['_VARIABLEMESSAGE']._serialized_end = 906 \ No newline at end of file diff --git a/philote_mdo/generated/data_pb2.pyi b/philote_mdo/generated/data_pb2.pyi index 500c23f..35791fc 100644 --- a/philote_mdo/generated/data_pb2.pyi +++ b/philote_mdo/generated/data_pb2.pyi @@ -50,6 +50,14 @@ class DisciplineProperties(_message.Message): def __init__(self, continuous: bool=..., differentiable: bool=..., provides_gradients: bool=..., name: _Optional[str]=..., version: _Optional[str]=...) -> None: ... +class JobHandle(_message.Message): + __slots__ = ('job_id',) + JOB_ID_FIELD_NUMBER: _ClassVar[int] + job_id: str + + def __init__(self, job_id: _Optional[str]=...) -> None: + ... + class StreamOptions(_message.Message): __slots__ = ('num_double',) NUM_DOUBLE_FIELD_NUMBER: _ClassVar[int] diff --git a/philote_mdo/generated/disciplines_pb2.py b/philote_mdo/generated/disciplines_pb2.py index 1425ba1..ae40b97 100644 --- a/philote_mdo/generated/disciplines_pb2.py +++ b/philote_mdo/generated/disciplines_pb2.py @@ -8,7 +8,7 @@ _sym_db = _symbol_database.Default() from google.protobuf import empty_pb2 as google_dot_protobuf_dot_empty__pb2 from . import data_pb2 as data__pb2 -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x11disciplines.proto\x12\x07philote\x1a\x1bgoogle/protobuf/empty.proto\x1a\ndata.proto2\xd0\x04\n\x11DisciplineService\x12B\n\x07GetInfo\x12\x16.google.protobuf.Empty\x1a\x1d.philote.DisciplineProperties"\x00\x12D\n\x10SetStreamOptions\x12\x16.philote.StreamOptions\x1a\x16.google.protobuf.Empty"\x00\x12E\n\x13GetAvailableOptions\x12\x16.google.protobuf.Empty\x1a\x14.philote.OptionsList"\x00\x12B\n\nSetOptions\x12\x1a.philote.DisciplineOptions\x1a\x16.google.protobuf.Empty"\x00\x129\n\x05Setup\x12\x16.google.protobuf.Empty\x1a\x16.google.protobuf.Empty"\x00\x12O\n\x16GetVariableDefinitions\x12\x16.google.protobuf.Empty\x1a\x19.philote.VariableMetaData"\x000\x01\x12N\n\x15GetPartialDefinitions\x12\x16.google.protobuf.Empty\x1a\x19.philote.PartialsMetaData"\x000\x01\x12J\n\x11SetVariableShapes\x12\x19.philote.VariableMetaData\x1a\x16.google.protobuf.Empty"\x00(\x012\xab\x01\n\x0fExplicitService\x12K\n\x0fComputeFunction\x12\x18.philote.VariableMessage\x1a\x18.philote.VariableMessage"\x00(\x010\x01\x12K\n\x0fComputeGradient\x12\x18.philote.VariableMessage\x1a\x18.philote.VariableMessage"\x00(\x010\x012\x81\x02\n\x0fImplicitService\x12L\n\x10ComputeResiduals\x12\x18.philote.VariableMessage\x1a\x18.philote.VariableMessage"\x00(\x010\x01\x12J\n\x0eSolveResiduals\x12\x18.philote.VariableMessage\x1a\x18.philote.VariableMessage"\x00(\x010\x01\x12T\n\x18ComputeResidualGradients\x12\x18.philote.VariableMessage\x1a\x18.philote.VariableMessage"\x00(\x010\x01B\x11\n\x0forg.philote.mdob\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x11disciplines.proto\x12\x07philote\x1a\x1bgoogle/protobuf/empty.proto\x1a\ndata.proto2\x85\x06\n\x11DisciplineService\x128\n\x08StartJob\x12\x16.google.protobuf.Empty\x1a\x12.philote.JobHandle"\x00\x12:\n\x06EndJob\x12\x16.google.protobuf.Empty\x1a\x16.google.protobuf.Empty"\x00\x12=\n\tKeepAlive\x12\x16.google.protobuf.Empty\x1a\x16.google.protobuf.Empty"\x00\x12B\n\x07GetInfo\x12\x16.google.protobuf.Empty\x1a\x1d.philote.DisciplineProperties"\x00\x12D\n\x10SetStreamOptions\x12\x16.philote.StreamOptions\x1a\x16.google.protobuf.Empty"\x00\x12E\n\x13GetAvailableOptions\x12\x16.google.protobuf.Empty\x1a\x14.philote.OptionsList"\x00\x12B\n\nSetOptions\x12\x1a.philote.DisciplineOptions\x1a\x16.google.protobuf.Empty"\x00\x129\n\x05Setup\x12\x16.google.protobuf.Empty\x1a\x16.google.protobuf.Empty"\x00\x12O\n\x16GetVariableDefinitions\x12\x16.google.protobuf.Empty\x1a\x19.philote.VariableMetaData"\x000\x01\x12N\n\x15GetPartialDefinitions\x12\x16.google.protobuf.Empty\x1a\x19.philote.PartialsMetaData"\x000\x01\x12J\n\x11SetVariableShapes\x12\x19.philote.VariableMetaData\x1a\x16.google.protobuf.Empty"\x00(\x012\xab\x01\n\x0fExplicitService\x12K\n\x0fComputeFunction\x12\x18.philote.VariableMessage\x1a\x18.philote.VariableMessage"\x00(\x010\x01\x12K\n\x0fComputeGradient\x12\x18.philote.VariableMessage\x1a\x18.philote.VariableMessage"\x00(\x010\x012\x81\x02\n\x0fImplicitService\x12L\n\x10ComputeResiduals\x12\x18.philote.VariableMessage\x1a\x18.philote.VariableMessage"\x00(\x010\x01\x12J\n\x0eSolveResiduals\x12\x18.philote.VariableMessage\x1a\x18.philote.VariableMessage"\x00(\x010\x01\x12T\n\x18ComputeResidualGradients\x12\x18.philote.VariableMessage\x1a\x18.philote.VariableMessage"\x00(\x010\x01B\x11\n\x0forg.philote.mdob\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) _builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'disciplines_pb2', _globals) @@ -16,8 +16,8 @@ _globals['DESCRIPTOR']._loaded_options = None _globals['DESCRIPTOR']._serialized_options = b'\n\x0forg.philote.mdo' _globals['_DISCIPLINESERVICE']._serialized_start = 72 - _globals['_DISCIPLINESERVICE']._serialized_end = 664 - _globals['_EXPLICITSERVICE']._serialized_start = 667 - _globals['_EXPLICITSERVICE']._serialized_end = 838 - _globals['_IMPLICITSERVICE']._serialized_start = 841 - _globals['_IMPLICITSERVICE']._serialized_end = 1098 \ No newline at end of file + _globals['_DISCIPLINESERVICE']._serialized_end = 845 + _globals['_EXPLICITSERVICE']._serialized_start = 848 + _globals['_EXPLICITSERVICE']._serialized_end = 1019 + _globals['_IMPLICITSERVICE']._serialized_start = 1022 + _globals['_IMPLICITSERVICE']._serialized_end = 1279 \ No newline at end of file diff --git a/philote_mdo/generated/disciplines_pb2_grpc.py b/philote_mdo/generated/disciplines_pb2_grpc.py index 638bd69..eb60294 100644 --- a/philote_mdo/generated/disciplines_pb2_grpc.py +++ b/philote_mdo/generated/disciplines_pb2_grpc.py @@ -27,6 +27,9 @@ def __init__(self, channel): Args: channel: A grpc.Channel. """ + self.StartJob = channel.unary_unary('/philote.DisciplineService/StartJob', request_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString, response_deserializer=data__pb2.JobHandle.FromString, _registered_method=True) + self.EndJob = channel.unary_unary('/philote.DisciplineService/EndJob', request_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString, response_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, _registered_method=True) + self.KeepAlive = channel.unary_unary('/philote.DisciplineService/KeepAlive', request_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString, response_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, _registered_method=True) self.GetInfo = channel.unary_unary('/philote.DisciplineService/GetInfo', request_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString, response_deserializer=data__pb2.DisciplineProperties.FromString, _registered_method=True) self.SetStreamOptions = channel.unary_unary('/philote.DisciplineService/SetStreamOptions', request_serializer=data__pb2.StreamOptions.SerializeToString, response_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, _registered_method=True) self.GetAvailableOptions = channel.unary_unary('/philote.DisciplineService/GetAvailableOptions', request_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString, response_deserializer=data__pb2.OptionsList.FromString, _registered_method=True) @@ -43,6 +46,38 @@ class DisciplineServiceServicer(object): bindings for implicit and explicit disciplines. """ + def StartJob(self, request, context): + """Starts a job and returns its handle. + + A job owns one discipline instance and all the state built on it. The + client must call this before any RPC other than GetInfo and + GetAvailableOptions, and must send the returned id in the + "philote-job-id" metadata header on every subsequent call. A call that + omits the header is rejected with FAILED_PRECONDITION; one that names an + unknown or expired job is rejected with NOT_FOUND. + + Servers may cap the number of concurrent jobs and reject further + requests with RESOURCE_EXHAUSTED. + """ + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details('Method not implemented!') + raise NotImplementedError('Method not implemented!') + + def EndJob(self, request, context): + """Ends the job named in the metadata header and releases its resources. + """ + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details('Method not implemented!') + raise NotImplementedError('Method not implemented!') + + def KeepAlive(self, request, context): + """Marks the job named in the metadata header as still in use, so that it + is not evicted while the client is idle between evaluations. + """ + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details('Method not implemented!') + raise NotImplementedError('Method not implemented!') + def GetInfo(self, request, context): """Gets the fundamental properties of the discipline """ @@ -101,7 +136,7 @@ def SetVariableShapes(self, request_iterator, context): raise NotImplementedError('Method not implemented!') def add_DisciplineServiceServicer_to_server(servicer, server): - rpc_method_handlers = {'GetInfo': grpc.unary_unary_rpc_method_handler(servicer.GetInfo, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=data__pb2.DisciplineProperties.SerializeToString), 'SetStreamOptions': grpc.unary_unary_rpc_method_handler(servicer.SetStreamOptions, request_deserializer=data__pb2.StreamOptions.FromString, response_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString), 'GetAvailableOptions': grpc.unary_unary_rpc_method_handler(servicer.GetAvailableOptions, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=data__pb2.OptionsList.SerializeToString), 'SetOptions': grpc.unary_unary_rpc_method_handler(servicer.SetOptions, request_deserializer=data__pb2.DisciplineOptions.FromString, response_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString), 'Setup': grpc.unary_unary_rpc_method_handler(servicer.Setup, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString), 'GetVariableDefinitions': grpc.unary_stream_rpc_method_handler(servicer.GetVariableDefinitions, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=data__pb2.VariableMetaData.SerializeToString), 'GetPartialDefinitions': grpc.unary_stream_rpc_method_handler(servicer.GetPartialDefinitions, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=data__pb2.PartialsMetaData.SerializeToString), 'SetVariableShapes': grpc.stream_unary_rpc_method_handler(servicer.SetVariableShapes, request_deserializer=data__pb2.VariableMetaData.FromString, response_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString)} + rpc_method_handlers = {'StartJob': grpc.unary_unary_rpc_method_handler(servicer.StartJob, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=data__pb2.JobHandle.SerializeToString), 'EndJob': grpc.unary_unary_rpc_method_handler(servicer.EndJob, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString), 'KeepAlive': grpc.unary_unary_rpc_method_handler(servicer.KeepAlive, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString), 'GetInfo': grpc.unary_unary_rpc_method_handler(servicer.GetInfo, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=data__pb2.DisciplineProperties.SerializeToString), 'SetStreamOptions': grpc.unary_unary_rpc_method_handler(servicer.SetStreamOptions, request_deserializer=data__pb2.StreamOptions.FromString, response_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString), 'GetAvailableOptions': grpc.unary_unary_rpc_method_handler(servicer.GetAvailableOptions, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=data__pb2.OptionsList.SerializeToString), 'SetOptions': grpc.unary_unary_rpc_method_handler(servicer.SetOptions, request_deserializer=data__pb2.DisciplineOptions.FromString, response_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString), 'Setup': grpc.unary_unary_rpc_method_handler(servicer.Setup, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString), 'GetVariableDefinitions': grpc.unary_stream_rpc_method_handler(servicer.GetVariableDefinitions, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=data__pb2.VariableMetaData.SerializeToString), 'GetPartialDefinitions': grpc.unary_stream_rpc_method_handler(servicer.GetPartialDefinitions, request_deserializer=google_dot_protobuf_dot_empty__pb2.Empty.FromString, response_serializer=data__pb2.PartialsMetaData.SerializeToString), 'SetVariableShapes': grpc.stream_unary_rpc_method_handler(servicer.SetVariableShapes, request_deserializer=data__pb2.VariableMetaData.FromString, response_serializer=google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString)} generic_handler = grpc.method_handlers_generic_handler('philote.DisciplineService', rpc_method_handlers) server.add_generic_rpc_handlers((generic_handler,)) server.add_registered_method_handlers('philote.DisciplineService', rpc_method_handlers) @@ -113,6 +148,18 @@ class DisciplineService(object): bindings for implicit and explicit disciplines. """ + @staticmethod + def StartJob(request, target, options=(), channel_credentials=None, call_credentials=None, insecure=False, compression=None, wait_for_ready=None, timeout=None, metadata=None): + return grpc.experimental.unary_unary(request, target, '/philote.DisciplineService/StartJob', google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString, data__pb2.JobHandle.FromString, options, channel_credentials, insecure, call_credentials, compression, wait_for_ready, timeout, metadata, _registered_method=True) + + @staticmethod + def EndJob(request, target, options=(), channel_credentials=None, call_credentials=None, insecure=False, compression=None, wait_for_ready=None, timeout=None, metadata=None): + return grpc.experimental.unary_unary(request, target, '/philote.DisciplineService/EndJob', google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString, google_dot_protobuf_dot_empty__pb2.Empty.FromString, options, channel_credentials, insecure, call_credentials, compression, wait_for_ready, timeout, metadata, _registered_method=True) + + @staticmethod + def KeepAlive(request, target, options=(), channel_credentials=None, call_credentials=None, insecure=False, compression=None, wait_for_ready=None, timeout=None, metadata=None): + return grpc.experimental.unary_unary(request, target, '/philote.DisciplineService/KeepAlive', google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString, google_dot_protobuf_dot_empty__pb2.Empty.FromString, options, channel_credentials, insecure, call_credentials, compression, wait_for_ready, timeout, metadata, _registered_method=True) + @staticmethod def GetInfo(request, target, options=(), channel_credentials=None, call_credentials=None, insecure=False, compression=None, wait_for_ready=None, timeout=None, metadata=None): return grpc.experimental.unary_unary(request, target, '/philote.DisciplineService/GetInfo', google_dot_protobuf_dot_empty__pb2.Empty.SerializeToString, data__pb2.DisciplineProperties.FromString, options, channel_credentials, insecure, call_credentials, compression, wait_for_ready, timeout, metadata, _registered_method=True) diff --git a/philote_mdo/utils/__init__.py b/philote_mdo/utils/__init__.py index 7a4dec0..a3c7f08 100644 --- a/philote_mdo/utils/__init__.py +++ b/philote_mdo/utils/__init__.py @@ -46,6 +46,10 @@ PhiloteError, PhiloteValidationError, PhiloteServerError, + PhiloteJobError, + JobNotFoundError, + JobCapacityError, + JobStateError, validate_name, validate_shape, validate_units, diff --git a/philote_mdo/utils/validation.py b/philote_mdo/utils/validation.py index 75395f2..dcfcc7d 100644 --- a/philote_mdo/utils/validation.py +++ b/philote_mdo/utils/validation.py @@ -61,6 +61,52 @@ class PhiloteServerError(PhiloteError, RuntimeError): pass +class PhiloteJobError(PhiloteError, RuntimeError): + """Base class for errors concerning a server-side job. + + Raised on the client when the server rejects a call because of the job + it names, and on the server by ``JobStore`` before the failure is + mapped onto a gRPC status code. + """ + + pass + + +class JobNotFoundError(PhiloteJobError): + """Raised when a job id is unknown to the server, or has expired. + + Maps to ``grpc.StatusCode.NOT_FOUND``. Clients must treat this as + terminal rather than silently starting a replacement job: the state + built up on the old job is gone, and an optimizer that carried on + against a fresh discipline would return plausible but wrong results. + """ + + pass + + +class JobCapacityError(PhiloteJobError): + """Raised when a server already holds its maximum number of jobs. + + Maps to ``grpc.StatusCode.RESOURCE_EXHAUSTED``. A job can own a mesh or + an initialised solver, so the limit is refused explicitly rather than + being allowed to exhaust memory. + """ + + pass + + +class JobStateError(PhiloteJobError): + """Raised when an RPC arrives out of order for the job's current state. + + Maps to ``grpc.StatusCode.FAILED_PRECONDITION``. Setting options after + ``Setup`` is the motivating case: the variable metadata has already been + built from the previous values, so accepting the call would leave the + job describing itself inconsistently. + """ + + pass + + # --------------------------------------------------------------------------- # Validation helpers # --------------------------------------------------------------------------- diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..f3bf9be --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,237 @@ +# Philote-Python +# +# Copyright 2022-2025 Christopher A. Lupp +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# +# This work has been cleared for public release, distribution unlimited, case +# number: AFRL-2023-5713. +# +# The views expressed are those of the authors and do not reflect the +# official guidance or position of the United States Government, the +# Department of Defense or of the United States Air Force. +# +# Statement from DoD: The Appearance of external hyperlinks does not +# constitute endorsement by the United States Department of Defense (DoD) of +# the linked websites, of the information, products, or services contained +# therein. The DoD does not exercise any editorial, security, or other +# control over the information you may find at these locations. +""" +Shared helpers for the test suite. + +Two things changed when servers gained per-job state, and between them they +account for nearly every test that had to be touched. + +First, RPC handlers now read the ``philote-job-id`` metadata header, so a +handler called directly needs a context that can produce one. A bare +``Mock()`` cannot: ``invocation_metadata()`` returns another ``Mock``, which is +not iterable. :func:`job_context` supplies one that works. + +Second, a server no longer owns a discipline. It builds one per job, so a test +that used to configure ``server._discipline`` now configures +``job.discipline``. :func:`make_server` sets both up in one call. +""" +from unittest.mock import Mock, patch + +from philote_mdo.general.job import JOB_METADATA_KEY + + +def job_context(job=None, job_id=None): + """ + Returns a fake ``ServicerContext`` naming a job. + + ``abort`` is left as a plain ``Mock``, so it records the call and returns + instead of raising the way real gRPC does. The suite depends on that: tests + assert ``context.abort.assert_called_once()`` and then go on to inspect + what the handler produced. + + Parameters + ---------- + job : Job, optional + Job to name. Takes precedence over ``job_id``. + job_id : str, optional + Job id to name, when there is no job object to hand. + + Returns + ------- + unittest.mock.Mock + Context whose ``invocation_metadata()`` carries the job header. + """ + if job is not None: + job_id = job.job_id + + context = Mock() + + if job_id is None: + context.invocation_metadata.return_value = () + else: + context.invocation_metadata.return_value = ((JOB_METADATA_KEY, job_id),) + + return context + + +def make_server(server_cls, discipline_factory, **kwargs): + """ + Builds a server, starts one job on it, and returns a context for that job. + + Parameters + ---------- + server_cls : type + ``DisciplineServer`` or one of its subclasses. + discipline_factory : callable + Zero-argument callable returning a fresh discipline. + **kwargs + Passed through to the server constructor. ``ttl`` defaults to ``None`` + so that no sweeper thread is started for a unit test. + + Returns + ------- + tuple + ``(server, job, context)``. Configure the discipline through + ``job.discipline``. + """ + kwargs.setdefault("ttl", None) + + server = server_cls(discipline=discipline_factory, **kwargs) + job = server._jobs.create() + + return server, job, job_context(job=job) + + +def bind_job(client, job_id="test-job"): + """ + Binds a client to a job id without making an RPC. + + Client-side unit tests patch only the service stub they are exercising, so + the real ``DisciplineServiceStub`` sits on an intercepted ``Mock`` channel. + Letting the lazy start fire there would attempt a genuine ``StartJob``. + Presetting the id keeps those tests on the subject they are testing. + + Parameters + ---------- + client : DisciplineClient + The client to bind. + job_id : str, optional + Id to bind to. + + Returns + ------- + DisciplineClient + The same client, for chaining. + """ + client._job_id = job_id + return client + + +def patch_discipline_stub(job_id="test-job"): + """ + Replaces ``DisciplineServiceStub`` with a mock for client-side unit tests. + + Two things make this necessary. Clients now wrap their channel in an + interceptor, so a ``Mock`` channel is no longer inert: gRPC's real + machinery runs on top of it and fails when it tries to unpack a mock + response. And components that claim a job from inside their constructor -- + the OpenMDAO bindings call ``send_options`` in ``__init__`` -- leave no + window to bind an id afterwards. + + Patching the base stub covers every call that goes through it, including + the lazy ``StartJob``. + + Parameters + ---------- + job_id : str, optional + Id that ``StartJob`` appears to return. + + Returns + ------- + unittest.mock._patch + Start it with ``.start()`` and stop it with ``.stop()``, or use it as + a context manager. + """ + patcher = patch( + "philote_mdo.generated.disciplines_pb2_grpc.DisciplineServiceStub" + ) + real_start = patcher.start + real_stop = patcher.stop + + def start(): + stub_cls = real_start() + stub_cls.return_value.StartJob.return_value = Mock(job_id=job_id) + return stub_cls + + patcher.start = start + patcher.stop = real_stop + + return patcher + + +def make_server_from_instance(server_cls, discipline, **kwargs): + """ + Builds a server whose factory always hands back one given instance. + + Isolation between jobs is the point of the factory, so production code + should never do this. It is useful in a unit test that configures a + discipline by hand and then calls the job-independent RPCs -- ``GetInfo`` + and ``GetAvailableOptions`` answer from an instance of the server's own, + and this makes that instance the same object the test configured. + + Parameters + ---------- + server_cls : type + ``DisciplineServer`` or one of its subclasses. + discipline : Discipline + The instance every job and every describe() call will receive. + **kwargs + Passed through to the server constructor. + + Returns + ------- + tuple + ``(server, job, context)``. + """ + return make_server(server_cls, lambda: discipline, **kwargs) + + +class Aborted(Exception): + """ + Stands in for the exception real gRPC raises out of ``context.abort``. + """ + + +def aborting_job_context(job=None, job_id=None): + """ + Returns a context whose ``abort`` raises, as real gRPC's does. + + :func:`job_context` leaves ``abort`` as a plain mock because most of the + suite asserts on it and then carries on inspecting the handler's output. + That is the wrong shape for testing an abort path itself: without an + exception the handler runs on past the abort and fails a second time + further down, so the status code under test gets buried. + + Parameters + ---------- + job : Job, optional + Job to name. + job_id : str, optional + Job id to name. + + Returns + ------- + unittest.mock.Mock + Context that raises :class:`Aborted` when the handler aborts. + """ + context = job_context(job=job, job_id=job_id) + context.abort.side_effect = Aborted + + return context diff --git a/tests/test_discipline_client.py b/tests/test_discipline_client.py index d5184bc..406f940 100644 --- a/tests/test_discipline_client.py +++ b/tests/test_discipline_client.py @@ -28,6 +28,8 @@ # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. import unittest + +from conftest import bind_job from unittest.mock import Mock, MagicMock, patch import numpy as np from google.protobuf.empty_pb2 import Empty @@ -49,7 +51,7 @@ def test_init(self): mock_channel = Mock() # Create an instance of YourClass with the mock channel - instance = DisciplineClient(mock_channel) + instance = bind_job(DisciplineClient(mock_channel)) # Assert that the attributes are initialized correctly self.assertTrue(instance.verbose) @@ -78,7 +80,7 @@ def test_get_discipline_info(self, mock_discipline_stub): version="1.2.3", ) - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client.get_discipline_info() # check the values of the response @@ -96,7 +98,7 @@ def test_send_stream_options(self, mock_discipline_stub): mock_channel = Mock() mock_stub = mock_discipline_stub.return_value - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) expected_num_double = 10 client._stream_options = expected_options = data.StreamOptions( num_double=expected_num_double, @@ -208,7 +210,7 @@ def test_run_setup(self, mock_discipline_stub): """ mock_channel = Mock() mock_stub = mock_discipline_stub.return_value - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client.run_setup() # assert that the 'setup' and 'setup_partials' methods were called @@ -222,7 +224,7 @@ def test_get_variable_definitions(self, mock_discipline_stub): """ mock_channel = Mock() mock_stub = mock_discipline_stub.return_value - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) input_definition = data.VariableMetaData( name="x", shape=[2, 2], units="m", type=data.VariableType.kInput @@ -268,7 +270,7 @@ def test_get_partial_definitions(self, mock_discipline_stub): """ mock_channel = Mock() mock_stub = mock_discipline_stub.return_value - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) partials_metadata = [ data.PartialsMetaData(name="input1", subname="output1"), @@ -348,7 +350,7 @@ def test_assemble_input_messages(self): Tests the _assemble_input_messages function of the Discipline Client. """ mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client._stream_options.num_double = 2 input_data = { @@ -397,7 +399,7 @@ def test_recover_outputs(self): Tests the _recover_outputs function of the Discipline Client. """ mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(2, 2)), @@ -439,7 +441,7 @@ def test_recover_residuals(self): Tests the _recover_residuals function of the Discipline Client. """ mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kResidual, shape=(2, 2)), @@ -481,7 +483,7 @@ def test_recover_partials(self): Tests the _recover_partials function of the Discipline Client. """ mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(1,)), @@ -581,7 +583,7 @@ def test_recover_outputs_empty_array_raises_error(self): Tests that _recover_outputs raises ValueError when array data is empty. """ mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(2,)), @@ -605,7 +607,7 @@ def test_recover_residuals_empty_array_raises_error(self): Tests that _recover_residuals raises ValueError when array data is empty. """ mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kResidual, shape=(2,)), @@ -629,7 +631,7 @@ def test_recover_partials_empty_array_raises_error(self): Tests that _recover_partials raises ValueError when array data is empty. """ mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(2,)), @@ -655,19 +657,19 @@ def test_recover_partials_empty_array_raises_error(self): def test_send_options_non_dict_raises(self): mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) with self.assertRaises(PhiloteValidationError): client.send_options("not a dict") def test_assemble_input_messages_non_dict_raises(self): mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) with self.assertRaises(PhiloteValidationError): client._assemble_input_messages("not a dict") def test_assemble_input_messages_non_array_value_raises(self): mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) with self.assertRaises(PhiloteValidationError): client._assemble_input_messages({"x": [1.0, 2.0]}) diff --git a/tests/test_discipline_server.py b/tests/test_discipline_server.py index 5c49ae3..0015a71 100644 --- a/tests/test_discipline_server.py +++ b/tests/test_discipline_server.py @@ -28,6 +28,8 @@ # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. import unittest + +from conftest import job_context, make_server, make_server_from_instance from unittest.mock import Mock import grpc @@ -50,16 +52,21 @@ def test_get_info(self): """ Tests the GetInfo RPC of the Discipline Server. """ - server = DisciplineServer() - server._discipline = Discipline() - server._discipline._is_continuous = True - server._discipline._is_differentiable = True - server._discipline._provides_gradients = True - server._discipline._name = "TestDiscipline" - server._discipline._version = "1.2.3" + discipline = Discipline() + discipline._is_continuous = True + discipline._is_differentiable = True + discipline._provides_gradients = True + discipline._name = "TestDiscipline" + discipline._version = "1.2.3" + + # GetInfo answers from an instance of the server's own, so point the + # factory at the one this test configured + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) # mock arguments - context = Mock() + context = job_context(job=job) request = Empty() # GetInfo is a unary RPC, so it must return a message (not a generator) @@ -78,28 +85,32 @@ def test_set_stream_options(self): """ Tests the SetStreamOptions RPC of the Discipline Server. """ - server = DisciplineServer() + discipline = Discipline() + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) # mock arguments - context = Mock() + context = job_context(job=job) request = data.StreamOptions(num_double=2) server.SetStreamOptions(request, context) # check that the streaming options were set properly - self.assertEqual(server._stream_opts.num_double, 2) + self.assertEqual(job.stream_opts.num_double, 2) def test_get_available_options(self): - server = DisciplineServer() + discipline = Discipline() + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) # mock the request and context parameters (since they are not used in this function) request_mock = Mock() - context_mock = None + context_mock = context - server._discipline = Discipline() - - # set the mock options_list to _discipline.options_list - server._discipline.options_list = { + # set the mock options_list to the discipline's options_list + discipline.options_list = { "option1": "bool", "option2": "int", "option3": "float", @@ -116,24 +127,27 @@ def test_get_available_options(self): self.assertEqual(results.type, expected_types) def test_set_options(self): - server = DisciplineServer() + discipline = Mock() + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) # mock the request and context parameters request_mock = Mock() - context_mock = Mock() + context_mock = job_context(job=job) # set some mock options in the request request_mock.options = {"key1": "value1", "key2": 42} # create a mock for the _discipline attribute discipline_mock = Mock() - server._discipline = discipline_mock + job.discipline = discipline_mock # call the SetOptions function with the mock parameters server.SetOptions(request_mock, context_mock) # assert that the discipline's initialize method was called with the expected options - server._discipline.set_options.assert_called_once_with( + job.discipline.set_options.assert_called_once_with( {"key1": "value1", "key2": 42} ) @@ -141,35 +155,38 @@ def test_setup(self): """ Tests the Setup RPC of the Discipline Server. """ - context = Mock() request = Empty() - server = DisciplineServer() + server, job, context = make_server_from_instance( + DisciplineServer, Mock() + ) - # mock the 'setup' and 'setup_partials' methods of 'self.discipline' - server._discipline = Mock() - server._discipline.setup.return_value = None - server._discipline.setup_partials.return_value = None + # mock the 'setup' and 'setup_partials' methods of the discipline + job.discipline.setup.return_value = None + job.discipline.setup_partials.return_value = None server.Setup(request, context) # assert that the 'setup' and 'setup_partials' methods were called - server._discipline.setup.assert_called_once() - server._discipline.setup_partials.assert_called_once() + job.discipline.setup.assert_called_once() + job.discipline.setup_partials.assert_called_once() def test_get_variable_definitions(self): """ Tests the GetVariableDefinitions RPC of the Discipline Server. """ - server = DisciplineServer() - server._discipline = Discipline() + discipline = Discipline() + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) + discipline = job.discipline # add an input and an output - server._discipline.add_input("x", shape=(2, 2), units="m") - server._discipline.add_output("f", shape=(1,), units="m**2") + job.discipline.add_input("x", shape=(2, 2), units="m") + job.discipline.add_output("f", shape=(1,), units="m**2") # mock arguments - context = Mock() + context = job_context(job=job) request = Empty() response_generator = server.GetVariableDefinitions(request, context) @@ -210,8 +227,8 @@ def test_preallocate_inputs_explicit(self): Tests the preallocation of inputs for the explicit discipline cas of the Discipline Server (outputs are not an input). """ - server = DisciplineServer() - discipline = server._discipline = Discipline() + server, job, context = make_server(DisciplineServer, Discipline) + discipline = job.discipline discipline.add_input("x", shape=(2, 2), units="m") discipline.add_input("y", shape=(3, 3, 3), units="m**2") discipline.add_output("f1", shape=(1,), units="m**3") @@ -223,7 +240,7 @@ def test_preallocate_inputs_explicit(self): outputs = {} flat_outputs = {} - server.preallocate_inputs(inputs, flat_inputs, outputs, flat_outputs) + server.preallocate_inputs(job, inputs, flat_inputs, outputs, flat_outputs) # check the number of inputs and outputs self.assertEqual(len(inputs), 2) @@ -247,8 +264,8 @@ def test_preallocate_inputs_implicit(self): Tests the preallocation of inputs for the implicit discipline cas of the Discipline Server (outputs are an input). """ - server = DisciplineServer() - discipline = server._discipline = Discipline() + server, job, context = make_server(DisciplineServer, Discipline) + discipline = job.discipline discipline.add_input("x", shape=(2, 2), units="m") discipline.add_input("y", shape=(3, 3, 3), units="m**2") discipline.add_output("f1", shape=(1,), units="m**3") @@ -260,7 +277,7 @@ def test_preallocate_inputs_implicit(self): outputs = {} flat_outputs = {} - server.preallocate_inputs(inputs, flat_inputs, outputs, flat_outputs) + server.preallocate_inputs(job, inputs, flat_inputs, outputs, flat_outputs) # check the number of inputs and outputs self.assertEqual(len(inputs), 2) @@ -303,8 +320,8 @@ def test_preallocate_partials(self): This test is designed to catch the edge cases where either f or x are scalar. """ - server = DisciplineServer() - discipline = server._discipline = Discipline() + server, job, context = make_server(DisciplineServer, Discipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="m") discipline.add_input("y", shape=(3, 3), units="m**2") discipline.add_output("f1", shape=(1,), units="m**3") @@ -315,7 +332,7 @@ def test_preallocate_partials(self): discipline.declare_partials("f2", "x") discipline.declare_partials("f2", "y") - jac = server.preallocate_partials() + jac = server.preallocate_partials(job) pairs = [("f1", "x"), ("f1", "y"), ("f2", "x"), ("f2", "y")] expected_shapes = [(1,), (3, 3), (2, 3), (2, 3, 3, 3)] @@ -331,9 +348,11 @@ def test_preallocate_partials_implicit_uses_the_residual_shape(self): the function shape must be resolved against the residual entry and not against the output that shares its name. """ - server = DisciplineServer() - discipline = server._discipline = Discipline() + discipline = Discipline() discipline._is_implicit = True + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) discipline.add_input("x", shape=(3,), units="m") discipline.add_output("y", shape=(2,), units="m") @@ -346,7 +365,7 @@ def test_preallocate_partials_implicit_uses_the_residual_shape(self): discipline.declare_partials("y", "x") discipline.declare_partials("y", "y") - jac = server.preallocate_partials() + jac = server.preallocate_partials(job) # d(residual y)/dx uses the residual (4,) and the input (3,) self.assertEqual(jac[("y", "x")].shape, (4, 3)) @@ -358,14 +377,16 @@ def test_preallocate_partials_unknown_variable(self): A partial declared against a variable that was never added reports a validation error rather than a bare KeyError. """ - server = DisciplineServer() - discipline = server._discipline = Discipline() + discipline = Discipline() + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) discipline.add_input("x", shape=(1,)) discipline.add_output("f", shape=(1,)) discipline.declare_partials("f", "missing") with self.assertRaises(PhiloteValidationError): - server.preallocate_partials() + server.preallocate_partials(job) def test_process_inputs(self): # create a mock request_iterator @@ -395,7 +416,7 @@ def test_process_inputs(self): ), ] - server = DisciplineServer() + server, job, context = make_server(DisciplineServer, Discipline) # create mock flat_inputs and flat_outputs dictionaries flat_inputs = {"x": np.zeros(6)} @@ -411,13 +432,15 @@ def test_get_available_options_with_dict_type(self): """ Tests that GetAvailableOptions correctly maps dict options to kStruct. """ - server = DisciplineServer() + discipline = Discipline() + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) request_mock = Mock() - context_mock = None + context_mock = context - server._discipline = Discipline() - server._discipline.options_list = { + discipline.options_list = { "config": "dict", "flag": "bool", } @@ -434,10 +457,10 @@ def test_set_options_with_nested_dict(self): """ Tests that SetOptions correctly passes nested dict values through. """ - server = DisciplineServer() + server, job, context = make_server(DisciplineServer, Discipline) request_mock = Mock() - context_mock = Mock() + context_mock = job_context(job=job) request_mock.options = { "config": {"solver": "newton", "tol": 1e-6, "nested": {"a": 1}}, @@ -445,11 +468,11 @@ def test_set_options_with_nested_dict(self): } discipline_mock = Mock() - server._discipline = discipline_mock + job.discipline = discipline_mock server.SetOptions(request_mock, context_mock) - server._discipline.set_options.assert_called_once_with( + job.discipline.set_options.assert_called_once_with( { "config": {"solver": "newton", "tol": 1e-6, "nested": {"a": 1}}, "name": "test", @@ -461,15 +484,17 @@ def test_get_available_options_invalid_type_aborts(self): Tests that GetAvailableOptions calls context.abort for invalid option types. """ - server = DisciplineServer() - discipline = server._discipline = Discipline() + discipline = Discipline() + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) # Add option with invalid type (bypasses add_option validation by # writing directly to options_list) discipline.options_list["invalid_option"] = "unknown_type" request = Empty() - context = Mock() + context = job_context(job=job) server.GetAvailableOptions(request, context) @@ -483,7 +508,7 @@ def test_process_inputs_empty_array_raises_error(self): Tests that process_inputs raises PhiloteValidationError when array data is empty. """ - server = DisciplineServer() + server, job, context = make_server(DisciplineServer, Discipline) # Create request with empty data array request_iterator = [ @@ -511,16 +536,17 @@ def test_get_available_options_general_exception_aborts(self): Tests that GetAvailableOptions calls context.abort with INTERNAL for unexpected exceptions. """ - server = DisciplineServer() discipline = Mock() # options_list property raises an unexpected error type(discipline).options_list = property( lambda self: (_ for _ in ()).throw(RuntimeError("unexpected")) ) - server._discipline = discipline + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) request = Mock() - context = Mock() + context = job_context(job=job) server.GetAvailableOptions(request, context) @@ -534,14 +560,14 @@ def test_set_options_validation_error_aborts(self): Tests that SetOptions calls context.abort with INVALID_ARGUMENT for PhiloteValidationError. """ - server = DisciplineServer() + server, job, context = make_server(DisciplineServer, Discipline) discipline = Mock() discipline.set_options.side_effect = PhiloteValidationError("bad option") - server._discipline = discipline + job.discipline = discipline request = Mock() request.options = {} - context = Mock() + context = job_context(job=job) server.SetOptions(request, context) @@ -555,14 +581,14 @@ def test_set_options_general_exception_aborts(self): Tests that SetOptions calls context.abort with INTERNAL for unexpected exceptions. """ - server = DisciplineServer() + server, job, context = make_server(DisciplineServer, Discipline) discipline = Mock() discipline.set_options.side_effect = RuntimeError("boom") - server._discipline = discipline + job.discipline = discipline request = Mock() request.options = {} - context = Mock() + context = job_context(job=job) server.SetOptions(request, context) @@ -576,13 +602,13 @@ def test_setup_validation_error_aborts(self): Tests that Setup calls context.abort with INVALID_ARGUMENT for PhiloteValidationError. """ - server = DisciplineServer() + server, job, context = make_server(DisciplineServer, Discipline) discipline = Mock() discipline.setup.side_effect = PhiloteValidationError("bad setup") - server._discipline = discipline + job.discipline = discipline request = Mock() - context = Mock() + context = job_context(job=job) server.Setup(request, context) @@ -596,13 +622,13 @@ def test_setup_general_exception_aborts(self): Tests that Setup calls context.abort with INTERNAL for unexpected exceptions. """ - server = DisciplineServer() + server, job, context = make_server(DisciplineServer, Discipline) discipline = Mock() discipline._clear_data.side_effect = RuntimeError("crash") - server._discipline = discipline + job.discipline = discipline request = Mock() - context = Mock() + context = job_context(job=job) server.Setup(request, context) diff --git a/tests/test_discrete_integration.py b/tests/test_discrete_integration.py index 79e1316..b31e985 100644 --- a/tests/test_discrete_integration.py +++ b/tests/test_discrete_integration.py @@ -108,8 +108,8 @@ class TestDiscreteIntegration(unittest.TestCase): def _start_server(self, discipline, port): """Helper to start a gRPC server with the given discipline.""" - server = grpc.server(futures.ThreadPoolExecutor(max_workers=4)) - explicit_server = pmdo.ExplicitServer(discipline=discipline) + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) + explicit_server = pmdo.ExplicitServer(discipline=lambda: discipline) explicit_server.attach_to_server(server) server.add_insecure_port(f"[::]:{port}") server.start() diff --git a/tests/test_discrete_variables.py b/tests/test_discrete_variables.py index 6b3811a..f2364a1 100644 --- a/tests/test_discrete_variables.py +++ b/tests/test_discrete_variables.py @@ -31,6 +31,9 @@ Unit tests for discrete variable support across the Philote stack. """ import unittest + +from philote_mdo.general import Discipline +from conftest import patch_discipline_stub, job_context, make_server, bind_job from unittest.mock import Mock, MagicMock, patch import numpy as np @@ -155,12 +158,12 @@ class TestDisciplineServerDiscrete(unittest.TestCase): def test_get_variable_definitions_includes_discrete(self): """GetVariableDefinitions should stream both continuous and discrete metadata.""" - server = DisciplineServer() - server._discipline = Discipline() - server._discipline.add_input("x", shape=(1,)) - server._discipline.add_discrete_input("mode") + server, job, context = make_server(DisciplineServer, Discipline) + discipline = job.discipline + job.discipline.add_input("x", shape=(1,)) + job.discipline.add_discrete_input("mode") - responses = list(server.GetVariableDefinitions(None, None)) + responses = list(server.GetVariableDefinitions(None, context)) self.assertEqual(len(responses), 2) types = [r.type for r in responses] self.assertIn(data.VariableType.kInput, types) @@ -168,7 +171,7 @@ def test_get_variable_definitions_includes_discrete(self): def test_process_inputs_with_discrete(self): """process_inputs should demux continuous and discrete messages.""" - server = DisciplineServer() + server, job, context = make_server(DisciplineServer, Discipline) flat_inputs = {"x": np.zeros(2)} discrete_inputs = {} @@ -229,7 +232,7 @@ def test_get_variable_definitions_separates_discrete(self, mock_stub_cls): ), ] - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client.get_variable_definitions() self.assertEqual(len(client._var_meta), 2) @@ -238,7 +241,7 @@ def test_get_variable_definitions_separates_discrete(self, mock_stub_cls): def test_assemble_input_messages_with_discrete(self): """_assemble_input_messages should include discrete messages.""" mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client._stream_options.num_double = 10 inputs = {"x": np.array([1.0])} @@ -263,7 +266,7 @@ def test_assemble_input_messages_with_discrete(self): def test_assemble_input_messages_with_discrete_outputs(self): """_assemble_input_messages should include discrete output messages.""" mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client._stream_options.num_double = 10 inputs = {"x": np.array([1.0])} @@ -287,7 +290,7 @@ def test_assemble_input_messages_with_discrete_outputs(self): def test_recover_outputs_with_discrete(self): """_recover_outputs should return (outputs, discrete_outputs) tuple.""" mock_channel = Mock() - client = DisciplineClient(mock_channel) + client = bind_job(DisciplineClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(1,)), ] @@ -324,14 +327,14 @@ class TestExplicitServerDiscrete(unittest.TestCase): def test_compute_function_with_discrete(self): """ComputeFunction should pass discrete data to discipline.compute.""" - server = ExplicitServer() + server, job, context = make_server(ExplicitServer, Discipline) discipline = ExplicitDiscipline() discipline.add_input("x", shape=(1,)) discipline.add_output("f", shape=(1,)) discipline.add_discrete_input("mode") discipline.add_discrete_output("status") - server._discipline = discipline - server._stream_opts.num_double = 10 + job.discipline = discipline + job.stream_opts.num_double = 10 captured = {} @@ -358,7 +361,7 @@ def compute(inputs, outputs, discrete_inputs, discrete_outputs): ), ] - responses = list(server.ComputeFunction(request_iterator, None)) + responses = list(server.ComputeFunction(request_iterator, context)) # Should have received the discrete input self.assertEqual(captured["discrete_inputs"]["mode"], "fast") @@ -379,14 +382,14 @@ def compute(inputs, outputs, discrete_inputs, discrete_outputs): def test_compute_gradient_with_discrete(self): """ComputeGradient should pass discrete data to compute_partials.""" - server = ExplicitServer() + server, job, context = make_server(ExplicitServer, Discipline) discipline = ExplicitDiscipline() discipline.add_input("x", shape=(1,)) discipline.add_output("f", shape=(1,)) discipline.add_discrete_input("mode") discipline.declare_partials("f", "x") - server._discipline = discipline - server._stream_opts.num_double = 10 + job.discipline = discipline + job.stream_opts.num_double = 10 captured = {} @@ -412,7 +415,7 @@ def compute_partials(inputs, jac, discrete_inputs): ), ] - responses = list(server.ComputeGradient(request_iterator, None)) + responses = list(server.ComputeGradient(request_iterator, context)) self.assertEqual(captured["discrete_inputs"]["mode"], "fast") self.assertEqual(len(responses), 1) @@ -457,10 +460,10 @@ def _make_request(self, x_val=1.0, f_val=0.0, mode_val="fast"): ] def test_compute_residuals_with_discrete(self): - server = ImplicitServer() + server, job, context = make_server(ImplicitServer, Discipline) discipline = self._make_discipline() - server._discipline = discipline - server._stream_opts.num_double = 10 + job.discipline = discipline + job.stream_opts.num_double = 10 captured = {} @@ -471,17 +474,17 @@ def compute_residuals(inputs, outputs, residuals, di, do): discipline.compute_residuals = compute_residuals responses = list( - server.ComputeResiduals(self._make_request(), None) + server.ComputeResiduals(self._make_request(), context) ) self.assertEqual(captured["mode"], "fast") self.assertGreater(len(responses), 0) def test_solve_residuals_with_discrete(self): - server = ImplicitServer() + server, job, context = make_server(ImplicitServer, Discipline) discipline = self._make_discipline() - server._discipline = discipline - server._stream_opts.num_double = 10 + job.discipline = discipline + job.stream_opts.num_double = 10 captured = {} @@ -492,17 +495,17 @@ def solve_residuals(inputs, outputs, di): discipline.solve_residuals = solve_residuals responses = list( - server.SolveResiduals(self._make_request(), None) + server.SolveResiduals(self._make_request(), context) ) self.assertEqual(captured["mode"], "fast") self.assertGreater(len(responses), 0) def test_compute_residual_gradients_with_discrete(self): - server = ImplicitServer() + server, job, context = make_server(ImplicitServer, Discipline) discipline = self._make_discipline() - server._discipline = discipline - server._stream_opts.num_double = 10 + job.discipline = discipline + job.stream_opts.num_double = 10 captured = {} @@ -513,7 +516,7 @@ def residual_partials(inputs, outputs, jac, di, do): discipline.residual_partials = residual_partials responses = list( - server.ComputeResidualGradients(self._make_request(), None) + server.ComputeResidualGradients(self._make_request(), context) ) self.assertEqual(captured["mode"], "fast") @@ -530,7 +533,7 @@ class TestExplicitClientDiscrete(unittest.TestCase): def test_run_compute_with_discrete(self, mock_stub_cls): mock_channel = Mock() mock_stub = mock_stub_cls.return_value - client = ExplicitClient(mock_channel) + client = bind_job(ExplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(1,)), ] @@ -560,7 +563,7 @@ def test_run_compute_with_discrete(self, mock_stub_cls): def test_run_compute_partials_with_discrete(self, mock_stub_cls): mock_channel = Mock() mock_stub = mock_stub_cls.return_value - client = ExplicitClient(mock_channel) + client = bind_job(ExplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(1,)), data.VariableMetaData(name="x", type=data.kInput, shape=(1,)), @@ -649,7 +652,11 @@ def test_compute_with_discrete_tuple_result(self, om_patch): from philote_mdo.openmdao import RemoteExplicitComponent mock_channel = Mock() - comp = RemoteExplicitComponent(channel=mock_channel) + + # the component claims a job inside __init__, so the stub has to be + # mocked around the construction itself + with patch_discipline_stub(): + comp = RemoteExplicitComponent(channel=mock_channel) client_mock = MagicMock() client_mock._var_meta = [ diff --git a/tests/test_dynamic_shapes.py b/tests/test_dynamic_shapes.py index 96cdeaf..4c84df0 100644 --- a/tests/test_dynamic_shapes.py +++ b/tests/test_dynamic_shapes.py @@ -29,6 +29,9 @@ # control over the information you may find at these locations. from concurrent import futures import unittest + +from philote_mdo.general import Discipline +from conftest import job_context, make_server, make_server_from_instance from unittest.mock import Mock import grpc @@ -127,18 +130,15 @@ class TestSetVariableShapesRPC(unittest.TestCase): """Unit tests for the SetVariableShapes RPC handler.""" def _make_server_with_dynamic_disc(self): - server = DisciplineServer() disc = Discipline() disc.add_input("x", dynamic_shape=True) disc.add_output("y", dynamic_shape=True) disc.add_input("z", shape=(2,)) # static - server._discipline = disc - return server + return make_server_from_instance(DisciplineServer, disc) def test_set_shapes_for_dynamic_variables(self): """SetVariableShapes updates shapes on dynamic variables.""" - server = self._make_server_with_dynamic_disc() - context = Mock() + server, job, context = self._make_server_with_dynamic_disc() x_meta = data.VariableMetaData( name="x", type=data.VariableType.kInput, shape=[5] @@ -149,7 +149,7 @@ def test_set_shapes_for_dynamic_variables(self): server.SetVariableShapes(iter([x_meta, y_meta]), context) # verify shapes were updated - for var in server._discipline._var_meta: + for var in job.discipline._var_meta: if var.name == "x": self.assertEqual(list(var.shape), [5]) if var.name == "y" and var.type == data.VariableType.kOutput: @@ -157,8 +157,7 @@ def test_set_shapes_for_dynamic_variables(self): def test_reject_shape_for_static_variable(self): """SetVariableShapes aborts when targeting a non-dynamic variable.""" - server = self._make_server_with_dynamic_disc() - context = Mock() + server, job, context = self._make_server_with_dynamic_disc() z_meta = data.VariableMetaData( name="z", type=data.VariableType.kInput, shape=[10] @@ -168,8 +167,7 @@ def test_reject_shape_for_static_variable(self): def test_reject_unknown_variable(self): """SetVariableShapes aborts when the variable name is not found.""" - server = self._make_server_with_dynamic_disc() - context = Mock() + server, job, context = self._make_server_with_dynamic_disc() meta = data.VariableMetaData( name="nope", type=data.VariableType.kInput, shape=[3] @@ -179,8 +177,7 @@ def test_reject_unknown_variable(self): def test_reject_invalid_shape(self): """SetVariableShapes aborts on invalid (non-positive) shape.""" - server = self._make_server_with_dynamic_disc() - context = Mock() + server, job, context = self._make_server_with_dynamic_disc() meta = data.VariableMetaData( name="x", type=data.VariableType.kInput, shape=[-1] @@ -190,10 +187,10 @@ def test_reject_invalid_shape(self): def test_preallocate_raises_when_shape_unset(self): """preallocate_inputs raises if a dynamic variable has no shape.""" - server = self._make_server_with_dynamic_disc() + server, job, context = self._make_server_with_dynamic_disc() with self.assertRaises(PhiloteValidationError): - server.preallocate_inputs({}, {}) + server.preallocate_inputs(job, {}, {}) # --------------------------------------------------------------- @@ -240,7 +237,7 @@ def test_flexible_compute(self): """Client sets shapes, then computes with the FlexibleDiscipline.""" server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = ExplicitServer(discipline=FlexibleDiscipline()) + discipline = ExplicitServer(discipline=FlexibleDiscipline) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -282,7 +279,7 @@ def test_flexible_compute_partials(self): """Client sets shapes, then computes partials.""" server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = ExplicitServer(discipline=FlexibleDiscipline()) + discipline = ExplicitServer(discipline=FlexibleDiscipline) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -319,7 +316,7 @@ def test_backward_compat_static_shapes(self): server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = ExplicitServer(discipline=Paraboloid()) + discipline = ExplicitServer(discipline=Paraboloid) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -376,7 +373,7 @@ def test_implicit_dynamic_shape_residual(self): """SetVariableShapes updates residual entries for implicit disciplines.""" server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = ImplicitServer(discipline=DynamicImplicit()) + discipline = ImplicitServer(discipline=DynamicImplicit) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -416,12 +413,12 @@ class TestSetVariableShapesGenericException(unittest.TestCase): """Tests the generic exception handler in SetVariableShapes.""" def test_generic_exception_aborts(self): - server = DisciplineServer() + server, job, context = make_server(DisciplineServer, Discipline) disc = Discipline() disc.add_input("x", dynamic_shape=True) - server._discipline = disc + job.discipline = disc - context = Mock() + context = job_context(job=job) # Craft an iterator that raises a non-validation exception def bad_iterator(): diff --git a/tests/test_edge_cases.py b/tests/test_edge_cases.py index 87d70d7..5bc3e19 100644 --- a/tests/test_edge_cases.py +++ b/tests/test_edge_cases.py @@ -28,6 +28,9 @@ # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. import unittest + +from philote_mdo.general import Discipline +from conftest import job_context, make_server, make_server_from_instance, bind_job from unittest.mock import Mock, MagicMock import grpc @@ -46,30 +49,32 @@ def test_attach_discipline(self): """ Test attaching a discipline to the server (line 62). """ - server = DisciplineServer() discipline = ExplicitDiscipline() # Test attach_discipline method - server.attach_discipline(discipline) + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) # Verify the discipline was attached - self.assertEqual(server._discipline, discipline) + self.assertEqual(job.discipline, discipline) def test_get_available_options_with_str_type(self): """ Test GetAvailableOptions with str option type (covers line 101). """ - server = DisciplineServer() discipline = Mock() # Mock the options_list attribute to return a dict with str type discipline.options_list = {"str_option": "str"} - server.attach_discipline(discipline) + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) # Create a mock request and context request = Mock() - context = Mock() + context = job_context(job=job) # This should work and exercise the str type branch result = server.GetAvailableOptions(request, context) @@ -81,15 +86,16 @@ def test_get_available_options_with_dict_type(self): """ Test GetAvailableOptions with dict option type (covers kStruct mapping). """ - server = DisciplineServer() discipline = Mock() discipline.options_list = {"config": "dict"} - server.attach_discipline(discipline) + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) request = Mock() - context = Mock() + context = job_context(job=job) result = server.GetAvailableOptions(request, context) @@ -102,17 +108,18 @@ def test_get_available_options_with_invalid_type(self): Test GetAvailableOptions with invalid option type aborts with INVALID_ARGUMENT. """ - server = DisciplineServer() discipline = Mock() # Mock the options_list attribute to return a dict with invalid type discipline.options_list = {"invalid_option": "invalid_type"} - server.attach_discipline(discipline) + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) # Create a mock request and context request = Mock() - context = Mock() + context = job_context(job=job) server.GetAvailableOptions(request, context) @@ -125,7 +132,6 @@ def test_process_inputs_with_empty_continuous_data(self): """ Test process_inputs with empty continuous data arrays. """ - server = DisciplineServer() discipline = Mock() # Set up discipline with continuous variables @@ -135,7 +141,9 @@ def test_process_inputs_with_empty_continuous_data(self): discipline._var_meta[0].shape = [2] discipline._var_meta[0].type = data.kInput - server.attach_discipline(discipline) + server, job, context = make_server_from_instance( + DisciplineServer, discipline + ) # Create a VariableMessage wrapping an Array with empty data message = data.VariableMessage( @@ -170,7 +178,7 @@ def test_recover_outputs_with_empty_data(self): """ # Create a mock channel channel = Mock() - client = DisciplineClient(channel) + client = bind_job(DisciplineClient(channel)) # Set up outputs structure client._var_meta = [Mock()] diff --git a/tests/test_explicit_client.py b/tests/test_explicit_client.py index 558ff01..1416a59 100644 --- a/tests/test_explicit_client.py +++ b/tests/test_explicit_client.py @@ -28,6 +28,8 @@ # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. import unittest + +from conftest import bind_job from unittest.mock import Mock, patch import grpc @@ -51,7 +53,7 @@ def test_compute(self, mock_explicit_stub): """ mock_channel = Mock() mock_stub = mock_explicit_stub.return_value - client = ExplicitClient(mock_channel) + client = bind_job(ExplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(3,)), data.VariableMetaData(name="x", type=data.kInput, shape=(2, 2)), @@ -96,7 +98,7 @@ def test_compute_partials(self, mock_explicit_stub): """ mock_channel = Mock() mock_stub = mock_explicit_stub.return_value - client = ExplicitClient(mock_channel) + client = bind_job(ExplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(1,)), data.VariableMetaData(name="x", type=data.kInput, shape=(2, 2)), @@ -140,7 +142,7 @@ def test_compute_partials(self, mock_explicit_stub): def test_run_compute_non_dict_raises(self): mock_channel = Mock() - client = ExplicitClient(mock_channel) + client = bind_job(ExplicitClient(mock_channel)) with self.assertRaises(PhiloteValidationError): client.run_compute("not a dict") @@ -148,7 +150,7 @@ def test_run_compute_non_dict_raises(self): def test_run_compute_grpc_error_wraps(self, mock_explicit_stub): mock_channel = Mock() mock_stub = mock_explicit_stub.return_value - client = ExplicitClient(mock_channel) + client = bind_job(ExplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="x", type=data.kInput, shape=(1,)), ] @@ -166,7 +168,7 @@ def test_run_compute_grpc_error_wraps(self, mock_explicit_stub): def test_run_compute_partials_grpc_error_wraps(self, mock_explicit_stub): mock_channel = Mock() mock_stub = mock_explicit_stub.return_value - client = ExplicitClient(mock_channel) + client = bind_job(ExplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(1,)), data.VariableMetaData(name="x", type=data.kInput, shape=(1,)), @@ -184,6 +186,6 @@ def test_run_compute_partials_grpc_error_wraps(self, mock_explicit_stub): def test_run_compute_partials_non_dict_raises(self): mock_channel = Mock() - client = ExplicitClient(mock_channel) + client = bind_job(ExplicitClient(mock_channel)) with self.assertRaises(PhiloteValidationError): client.run_compute_partials(42) diff --git a/tests/test_explicit_server.py b/tests/test_explicit_server.py index f4110b7..75199bc 100644 --- a/tests/test_explicit_server.py +++ b/tests/test_explicit_server.py @@ -28,6 +28,8 @@ # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. import unittest + +from conftest import job_context, make_server from unittest.mock import Mock import grpc @@ -50,13 +52,13 @@ def test_compute_function(self): """ Tests the ComputeFunction RPC of the Explicit Server. """ - server = ExplicitServer() - discipline = server._discipline = ExplicitDiscipline() - server._stream_opts.num_double = 3 + server, job, context = make_server(ExplicitServer, ExplicitDiscipline) + discipline = job.discipline + job.stream_opts.num_double = 3 discipline.add_input("x", shape=(5,), units="") discipline.add_output("f", shape=(2,), units="") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -76,7 +78,7 @@ def test_compute_function(self): def compute(inputs, outputs): outputs["f"] = np.array([rosen(inputs["x"]), rosen(inputs["x"].T - 2.0)]) - server._discipline.compute = compute + job.discipline.compute = compute # call the function response_generator = server.ComputeFunction(request_iterator, context) @@ -97,14 +99,14 @@ def test_compute_gradient(self): """ Tests the ComputeGradient RPC of the Explicit Server. """ - server = ExplicitServer() - discipline = server._discipline = ExplicitDiscipline() - server._stream_opts.num_double = 3 + server, job, context = make_server(ExplicitServer, ExplicitDiscipline) + discipline = job.discipline + job.stream_opts.num_double = 3 discipline.add_input("x", shape=(5,), units="") discipline.add_output("f", shape=(1,), units="") discipline.declare_partials("f", "x") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -124,7 +126,7 @@ def test_compute_gradient(self): def compute_partials(inputs, jac): jac["f", "x"] = rosen_der(inputs["x"]) - server._discipline.compute_partials = compute_partials + job.discipline.compute_partials = compute_partials # call the function response_generator = server.ComputeGradient(request_iterator, context) @@ -153,12 +155,12 @@ def test_compute_function_aborts_on_validation_error(self): Tests that ComputeFunction calls context.abort with INVALID_ARGUMENT when a PhiloteValidationError is raised. """ - server = ExplicitServer() - discipline = server._discipline = ExplicitDiscipline() + server, job, context = make_server(ExplicitServer, ExplicitDiscipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="") discipline.add_output("f", shape=(1,), units="") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -171,7 +173,7 @@ def test_compute_function_aborts_on_validation_error(self): def bad_compute(inputs, outputs): raise PhiloteValidationError("bad input data") - server._discipline.compute = bad_compute + job.discipline.compute = bad_compute list(server.ComputeFunction(request_iterator, context)) @@ -185,13 +187,13 @@ def test_compute_gradient_aborts_on_validation_error(self): Tests that ComputeGradient calls context.abort with INVALID_ARGUMENT when a PhiloteValidationError is raised. """ - server = ExplicitServer() - discipline = server._discipline = ExplicitDiscipline() + server, job, context = make_server(ExplicitServer, ExplicitDiscipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="") discipline.add_output("f", shape=(1,), units="") discipline.declare_partials("f", "x") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -204,7 +206,7 @@ def test_compute_gradient_aborts_on_validation_error(self): def bad_partials(inputs, jac): raise PhiloteValidationError("invalid partials") - server._discipline.compute_partials = bad_partials + job.discipline.compute_partials = bad_partials list(server.ComputeGradient(request_iterator, context)) @@ -218,12 +220,12 @@ def test_compute_function_aborts_on_discipline_error(self): Tests that ComputeFunction calls context.abort when the discipline's compute raises an exception. """ - server = ExplicitServer() - discipline = server._discipline = ExplicitDiscipline() + server, job, context = make_server(ExplicitServer, ExplicitDiscipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="") discipline.add_output("f", shape=(1,), units="") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -236,7 +238,7 @@ def test_compute_function_aborts_on_discipline_error(self): def bad_compute(inputs, outputs): raise RuntimeError("division by zero in compute") - server._discipline.compute = bad_compute + job.discipline.compute = bad_compute # Exhaust the generator list(server.ComputeFunction(request_iterator, context)) @@ -251,13 +253,13 @@ def test_compute_gradient_aborts_on_discipline_error(self): Tests that ComputeGradient calls context.abort when the discipline's compute_partials raises an exception. """ - server = ExplicitServer() - discipline = server._discipline = ExplicitDiscipline() + server, job, context = make_server(ExplicitServer, ExplicitDiscipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="") discipline.add_output("f", shape=(1,), units="") discipline.declare_partials("f", "x") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -270,7 +272,7 @@ def test_compute_gradient_aborts_on_discipline_error(self): def bad_partials(inputs, jac): raise RuntimeError("singular matrix") - server._discipline.compute_partials = bad_partials + job.discipline.compute_partials = bad_partials list(server.ComputeGradient(request_iterator, context)) diff --git a/tests/test_implicit_client.py b/tests/test_implicit_client.py index d9333c9..c1f95ed 100644 --- a/tests/test_implicit_client.py +++ b/tests/test_implicit_client.py @@ -28,6 +28,8 @@ # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. import unittest + +from conftest import bind_job from unittest.mock import Mock, patch import grpc @@ -51,7 +53,7 @@ def test_compute_residuals(self, mock_implicit_stub): """ mock_channel = Mock() mock_stub = mock_implicit_stub.return_value - client = ImplicitClient(mock_channel) + client = bind_job(ImplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(3,)), data.VariableMetaData(name="f", type=data.kResidual, shape=(3,)), @@ -103,7 +105,7 @@ def test_solve_residuals(self, mock_implicit_stub): """ mock_channel = Mock() mock_stub = mock_implicit_stub.return_value - client = ImplicitClient(mock_channel) + client = bind_job(ImplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(3,)), data.VariableMetaData(name="f", type=data.kResidual, shape=(3,)), @@ -154,7 +156,7 @@ def test_residual_partials(self, mock_implicit_stub): """ mock_channel = Mock() mock_stub = mock_implicit_stub.return_value - client = ImplicitClient(mock_channel) + client = bind_job(ImplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="f", type=data.kOutput, shape=(2,)), data.VariableMetaData(name="f", type=data.kResidual, shape=(2,)), @@ -208,7 +210,7 @@ def test_residual_partials(self, mock_implicit_stub): def test_run_compute_residuals_grpc_error_wraps(self, mock_implicit_stub): mock_channel = Mock() mock_stub = mock_implicit_stub.return_value - client = ImplicitClient(mock_channel) + client = bind_job(ImplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="x", type=data.kInput, shape=(1,)), data.VariableMetaData(name="f", type=data.kOutput, shape=(1,)), @@ -230,7 +232,7 @@ def test_run_compute_residuals_grpc_error_wraps(self, mock_implicit_stub): def test_run_solve_residuals_grpc_error_wraps(self, mock_implicit_stub): mock_channel = Mock() mock_stub = mock_implicit_stub.return_value - client = ImplicitClient(mock_channel) + client = bind_job(ImplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="x", type=data.kInput, shape=(1,)), data.VariableMetaData(name="f", type=data.kOutput, shape=(1,)), @@ -249,7 +251,7 @@ def test_run_solve_residuals_grpc_error_wraps(self, mock_implicit_stub): def test_run_residual_gradients_grpc_error_wraps(self, mock_implicit_stub): mock_channel = Mock() mock_stub = mock_implicit_stub.return_value - client = ImplicitClient(mock_channel) + client = bind_job(ImplicitClient(mock_channel)) client._var_meta = [ data.VariableMetaData(name="x", type=data.kInput, shape=(1,)), data.VariableMetaData(name="f", type=data.kOutput, shape=(1,)), @@ -270,25 +272,25 @@ def test_run_residual_gradients_grpc_error_wraps(self, mock_implicit_stub): def test_run_compute_residuals_non_dict_inputs_raises(self): mock_channel = Mock() - client = ImplicitClient(mock_channel) + client = bind_job(ImplicitClient(mock_channel)) with self.assertRaises(PhiloteValidationError): client.run_compute_residuals("not a dict", {"f": np.array([1.0])}) def test_run_compute_residuals_non_dict_outputs_raises(self): mock_channel = Mock() - client = ImplicitClient(mock_channel) + client = bind_job(ImplicitClient(mock_channel)) with self.assertRaises(PhiloteValidationError): client.run_compute_residuals({"x": np.array([1.0])}, "not a dict") def test_run_solve_residuals_non_dict_raises(self): mock_channel = Mock() - client = ImplicitClient(mock_channel) + client = bind_job(ImplicitClient(mock_channel)) with self.assertRaises(PhiloteValidationError): client.run_solve_residuals(42) def test_run_residual_gradients_non_dict_raises(self): mock_channel = Mock() - client = ImplicitClient(mock_channel) + client = bind_job(ImplicitClient(mock_channel)) with self.assertRaises(PhiloteValidationError): client.run_residual_gradients("bad", {"f": np.array([1.0])}) diff --git a/tests/test_implicit_server.py b/tests/test_implicit_server.py index dd134b2..496d132 100644 --- a/tests/test_implicit_server.py +++ b/tests/test_implicit_server.py @@ -28,6 +28,9 @@ # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. import unittest + +from philote_mdo.general import Discipline +from conftest import job_context, make_server from unittest.mock import Mock import grpc @@ -52,9 +55,9 @@ def test_compute_residuals(self): compute_residuals function, so that the entire solution process is tested (the actual residual function is mocked). """ - server = ImplicitServer() - server._stream_opts.num_double = 1 - discipline = server._discipline = ImplicitDiscipline() + server, job, context = make_server(ImplicitServer, Discipline) + job.stream_opts.num_double = 1 + discipline = job.discipline = ImplicitDiscipline() discipline.add_input("x", shape=(2,), units="m") discipline.add_input("y", shape=(2,), units="m") discipline.add_output("f", shape=(2,), units="m") @@ -76,10 +79,10 @@ def test_compute_residuals(self): def compute_residuals(inputs, outputs, residuals): residuals["f"] = np.array([7.0, 8.0]) - server._discipline.compute_residuals = compute_residuals + job.discipline.compute_residuals = compute_residuals # call the ComputeResiduals method - response_generator = server.ComputeResiduals(mock_request_iterator, None) + response_generator = server.ComputeResiduals(mock_request_iterator, context) result = list(response_generator) # assert that the expected residual messages were yielded @@ -107,9 +110,9 @@ def test_solve_residuals(self): solve_residuals function, so that the entire solution process is tested (the actual residual function is mocked). """ - server = ImplicitServer() - server._stream_opts.num_double = 1 - discipline = server._discipline = ImplicitDiscipline() + server, job, context = make_server(ImplicitServer, Discipline) + job.stream_opts.num_double = 1 + discipline = job.discipline = ImplicitDiscipline() discipline.add_input("x", shape=(2,), units="m") discipline.add_input("y", shape=(2,), units="m") discipline.add_output("f", shape=(2,), units="m") @@ -131,10 +134,10 @@ def test_solve_residuals(self): def solve_residuals(inputs, outputs): outputs["f"] = np.array([7.0, 8.0]) - server._discipline.solve_residuals = solve_residuals + job.discipline.solve_residuals = solve_residuals # call the SolveResiduals method - response_generator = server.SolveResiduals(mock_request_iterator, None) + response_generator = server.SolveResiduals(mock_request_iterator, context) result = list(response_generator) # assert that the expected output messages were yielded @@ -158,14 +161,14 @@ def test_residual_gradients(self): """ Tests the ComputeResiduals RPC of the Implicit Server. """ - server = ImplicitServer() - discipline = server._discipline = ImplicitDiscipline() - server._stream_opts.num_double = 3 + server, job, context = make_server(ImplicitServer, ImplicitDiscipline) + discipline = job.discipline + job.stream_opts.num_double = 3 discipline.add_input("x", shape=(5,), units="") discipline.add_output("f", shape=(1,), units="") discipline.declare_partials("f", "x") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -185,7 +188,7 @@ def test_residual_gradients(self): def residual_partials(inputs, residuals, jac): jac["f", "x"] = np.array([-251.0, -499.0, 11105.0, 25007.0, -2950.0]) - server._discipline.residual_partials = residual_partials + job.discipline.residual_partials = residual_partials # call the function response_generator = server.ComputeResidualGradients(request_iterator, context) @@ -215,12 +218,12 @@ def test_compute_residuals_aborts_on_validation_error(self): Tests that ComputeResiduals calls context.abort with INVALID_ARGUMENT when a PhiloteValidationError is raised. """ - server = ImplicitServer() - discipline = server._discipline = ImplicitDiscipline() + server, job, context = make_server(ImplicitServer, ImplicitDiscipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="") discipline.add_output("f", shape=(1,), units="") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -239,7 +242,7 @@ def test_compute_residuals_aborts_on_validation_error(self): def bad_residuals(inputs, outputs, residuals): raise PhiloteValidationError("bad residual input") - server._discipline.compute_residuals = bad_residuals + job.discipline.compute_residuals = bad_residuals list(server.ComputeResiduals(request_iterator, context)) @@ -253,12 +256,12 @@ def test_solve_residuals_aborts_on_validation_error(self): Tests that SolveResiduals calls context.abort with INVALID_ARGUMENT when a PhiloteValidationError is raised. """ - server = ImplicitServer() - discipline = server._discipline = ImplicitDiscipline() + server, job, context = make_server(ImplicitServer, ImplicitDiscipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="") discipline.add_output("f", shape=(1,), units="") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -277,7 +280,7 @@ def test_solve_residuals_aborts_on_validation_error(self): def bad_solve(inputs, outputs): raise PhiloteValidationError("bad solve input") - server._discipline.solve_residuals = bad_solve + job.discipline.solve_residuals = bad_solve list(server.SolveResiduals(request_iterator, context)) @@ -291,13 +294,13 @@ def test_compute_residual_gradients_aborts_on_validation_error(self): Tests that ComputeResidualGradients calls context.abort with INVALID_ARGUMENT when a PhiloteValidationError is raised. """ - server = ImplicitServer() - discipline = server._discipline = ImplicitDiscipline() + server, job, context = make_server(ImplicitServer, ImplicitDiscipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="") discipline.add_output("f", shape=(1,), units="") discipline.declare_partials("f", "x") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -310,7 +313,7 @@ def test_compute_residual_gradients_aborts_on_validation_error(self): def bad_partials(inputs, outputs, jac): raise PhiloteValidationError("bad partials input") - server._discipline.residual_partials = bad_partials + job.discipline.residual_partials = bad_partials list(server.ComputeResidualGradients(request_iterator, context)) @@ -324,13 +327,13 @@ def test_compute_residual_gradients_aborts_on_discipline_error(self): Tests that ComputeResidualGradients calls context.abort with INTERNAL when an unexpected exception is raised. """ - server = ImplicitServer() - discipline = server._discipline = ImplicitDiscipline() + server, job, context = make_server(ImplicitServer, ImplicitDiscipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="") discipline.add_output("f", shape=(1,), units="") discipline.declare_partials("f", "x") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -343,7 +346,7 @@ def test_compute_residual_gradients_aborts_on_discipline_error(self): def bad_partials(inputs, outputs, jac): raise RuntimeError("unexpected crash") - server._discipline.residual_partials = bad_partials + job.discipline.residual_partials = bad_partials list(server.ComputeResidualGradients(request_iterator, context)) @@ -357,12 +360,12 @@ def test_compute_residuals_aborts_on_discipline_error(self): Tests that ComputeResiduals calls context.abort when the discipline's compute_residuals raises an exception. """ - server = ImplicitServer() - discipline = server._discipline = ImplicitDiscipline() + server, job, context = make_server(ImplicitServer, ImplicitDiscipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="") discipline.add_output("f", shape=(1,), units="") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -381,7 +384,7 @@ def test_compute_residuals_aborts_on_discipline_error(self): def bad_residuals(inputs, outputs, residuals): raise RuntimeError("residual computation failed") - server._discipline.compute_residuals = bad_residuals + job.discipline.compute_residuals = bad_residuals list(server.ComputeResiduals(request_iterator, context)) @@ -395,12 +398,12 @@ def test_solve_residuals_aborts_on_discipline_error(self): Tests that SolveResiduals calls context.abort when the discipline's solve_residuals raises an exception. """ - server = ImplicitServer() - discipline = server._discipline = ImplicitDiscipline() + server, job, context = make_server(ImplicitServer, ImplicitDiscipline) + discipline = job.discipline discipline.add_input("x", shape=(1,), units="") discipline.add_output("f", shape=(1,), units="") - context = Mock() + context = job_context(job=job) request_iterator = [ data.VariableMessage( continuous=data.Array( @@ -419,7 +422,7 @@ def test_solve_residuals_aborts_on_discipline_error(self): def bad_solve(inputs, outputs): raise RuntimeError("solver did not converge") - server._discipline.solve_residuals = bad_solve + job.discipline.solve_residuals = bad_solve list(server.SolveResiduals(request_iterator, context)) diff --git a/tests/test_integration.py b/tests/test_integration.py index 9cfe0c8..6b1a3c5 100644 --- a/tests/test_integration.py +++ b/tests/test_integration.py @@ -29,6 +29,8 @@ # control over the information you may find at these locations. from concurrent import futures import unittest + +from conftest import job_context, make_server import grpc import numpy as np import philote_mdo.general as pmdo @@ -48,7 +50,7 @@ def test_paraboloid_compute(self): # server code server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = pmdo.ExplicitServer(discipline=Paraboloid()) + discipline = pmdo.ExplicitServer(discipline=Paraboloid) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -83,7 +85,7 @@ def test_paraboloid_compute_partials(self): # server code server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = pmdo.ExplicitServer(discipline=Paraboloid()) + discipline = pmdo.ExplicitServer(discipline=Paraboloid) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -119,7 +121,7 @@ def test_quadratic_compute_residuals(self): # server code server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = pmdo.ImplicitServer(discipline=QuadradicImplicit()) + discipline = pmdo.ImplicitServer(discipline=QuadradicImplicit) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -155,7 +157,7 @@ def test_quadratic_solve_residuals(self): # server code server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = pmdo.ImplicitServer(discipline=QuadradicImplicit()) + discipline = pmdo.ImplicitServer(discipline=QuadradicImplicit) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -190,7 +192,7 @@ def test_quadratic_residual_gradients(self): # server code server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = pmdo.ImplicitServer(discipline=QuadradicImplicit()) + discipline = pmdo.ImplicitServer(discipline=QuadradicImplicit) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -248,7 +250,7 @@ def residual_partials(self, inputs, outputs, jacobian): # server code server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - pmdo.ImplicitServer(discipline=VectorImplicit()).attach_to_server(server) + pmdo.ImplicitServer(discipline=VectorImplicit).attach_to_server(server) server.add_insecure_port("[::]:50051") server.start() @@ -289,7 +291,9 @@ def test_get_discipline_info(self): discipline = Paraboloid() discipline._name = "Paraboloid" discipline._version = "1.0.0" - pmdo.ExplicitServer(discipline=discipline).attach_to_server(server) + pmdo.ExplicitServer( + discipline=lambda: discipline + ).attach_to_server(server) server.add_insecure_port("[::]:50051") server.start() @@ -356,7 +360,7 @@ def test_struct_option_round_trip(self): """ # server server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) - discipline = pmdo.ExplicitServer(discipline=StructOptionDiscipline()) + discipline = pmdo.ExplicitServer(discipline=StructOptionDiscipline) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") server.start() diff --git a/tests/test_jobs.py b/tests/test_jobs.py new file mode 100644 index 0000000..0828d51 --- /dev/null +++ b/tests/test_jobs.py @@ -0,0 +1,793 @@ +# Philote-Python +# +# Copyright 2022-2025 Christopher A. Lupp +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# +# This work has been cleared for public release, distribution unlimited, case +# number: AFRL-2023-5713. +# +# The views expressed are those of the authors and do not reflect the +# official guidance or position of the United States Government, the +# Department of Defense or of the United States Air Force. +# +# Statement from DoD: The Appearance of external hyperlinks does not +# constitute endorsement by the United States Department of Defense (DoD) of +# the linked websites, of the information, products, or services contained +# therein. The DoD does not exercise any editorial, security, or other +# control over the information you may find at these locations. +""" +Tests for per-job state isolation. + +The bug these exist to prevent: a server used to share one discipline instance +across every client, so a second client's ``Setup`` rewrote the variable +metadata the first was still using. +""" +import threading +import time +import unittest +from concurrent import futures +from unittest.mock import Mock + +import grpc +import numpy as np +from scipy.optimize import rosen + +import philote_mdo.general as pmdo +import philote_mdo.generated.data_pb2 as data +from philote_mdo.examples import Paraboloid, Rosenbrock +from philote_mdo.general import Discipline, ExplicitDiscipline +from philote_mdo.general.discipline_server import DisciplineServer +from philote_mdo.general.job import JOB_METADATA_KEY, JobState, JobStore +from philote_mdo.utils.validation import ( + JobCapacityError, + PhiloteServerError, + JobNotFoundError, + JobStateError, + PhiloteJobError, +) + +from conftest import Aborted, aborting_job_context, job_context, make_server + + +def serve(discipline_factory, server_cls=pmdo.ExplicitServer, workers=16, **kwargs): + """ + Starts a real gRPC server on an ephemeral port. + + Returns + ------- + tuple + ``(grpc_server, port)``. Stop the server with ``.stop(0)``. + """ + server = grpc.server(futures.ThreadPoolExecutor(max_workers=workers)) + server_cls(discipline=discipline_factory, **kwargs).attach_to_server(server) + port = server.add_insecure_port("[::]:0") + server.start() + + return server, port + + +class TestJobIsolation(unittest.TestCase): + """Two clients against one server must not affect each other.""" + + def test_concurrent_clients_with_different_shapes(self): + """ + The bug from issue #76, reproduced exactly. + + Rosenbrock takes its variable shape from an option, so before jobs + existed whichever client called SetOptions last fixed the shapes both + clients got, and the other silently received a zero-padded result. + """ + server, port = serve(Rosenbrock) + self.addCleanup(server.stop, 0) + + def client_for(dimension): + client = pmdo.ExplicitClient( + channel=grpc.insecure_channel(f"localhost:{port}") + ) + client.send_options({"dimension": dimension}) + return client + + a, b = client_for(2), client_for(10) + + # interleaved exactly as in the issue: A sets up and reads its + # metadata, then B sets up, then both compute + a.run_setup() + a.get_variable_definitions() + b.run_setup() + b.get_variable_definitions() + + xa = np.array([1.5, 0.5]) + xb = np.arange(1.0, 11.0) + + self.assertAlmostEqual(a.run_compute({"x": xa})["f"][0], rosen(xa)) + self.assertAlmostEqual(b.run_compute({"x": xb})["f"][0], rosen(xb)) + + # and A is still correct after B has been all the way through setup + self.assertAlmostEqual(a.run_compute({"x": xa})["f"][0], rosen(xa)) + + self.assertNotEqual(a.job_id, b.job_id) + + def test_options_do_not_leak_between_jobs(self): + """A job's options are invisible to another job.""" + server, port = serve(Rosenbrock) + self.addCleanup(server.stop, 0) + + a = pmdo.ExplicitClient(channel=grpc.insecure_channel(f"localhost:{port}")) + b = pmdo.ExplicitClient(channel=grpc.insecure_channel(f"localhost:{port}")) + + a.send_options({"dimension": 3}) + b.send_options({"dimension": 7}) + + a.run_setup() + a.get_variable_definitions() + b.run_setup() + b.get_variable_definitions() + + shape_of = lambda c: tuple( + v.shape for v in c._var_meta if v.name == "x" + )[0] + + self.assertEqual(list(shape_of(a)), [3]) + self.assertEqual(list(shape_of(b)), [7]) + + def test_jobs_evaluate_concurrently(self): + """ + Two jobs may be inside compute() at the same time. + + The discipline here releases the GIL, which is what a compiled solver + does. A pure-Python discipline is still serialised by the interpreter; + jobs buy correctness unconditionally and throughput conditionally. + """ + hold = 0.4 + + class Slow(ExplicitDiscipline): + def setup(self): + self.add_input("x", shape=(1,)) + self.add_output("f", shape=(1,)) + + def compute(self, inputs, outputs): + time.sleep(hold) + outputs["f"] = inputs["x"] * 2.0 + + server, port = serve(Slow) + self.addCleanup(server.stop, 0) + + def run(store, index): + client = pmdo.ExplicitClient( + channel=grpc.insecure_channel(f"localhost:{port}") + ) + client.run_setup() + client.get_variable_definitions() + store[index] = client.run_compute({"x": np.array([float(index)])})["f"] + + results = [None, None] + threads = [threading.Thread(target=run, args=(results, i)) for i in range(2)] + + start = time.perf_counter() + for t in threads: + t.start() + for t in threads: + t.join() + elapsed = time.perf_counter() - start + + self.assertAlmostEqual(results[0][0], 0.0) + self.assertAlmostEqual(results[1][0], 2.0) + + # serialised would take 2 * hold; overlapped takes about one + self.assertLess(elapsed, 2 * hold) + + +class TestJobLifecycle(unittest.TestCase): + """StartJob, EndJob, KeepAlive and the state machine.""" + + def test_client_starts_a_job_lazily(self): + """An unmodified client script acquires a job on its first call.""" + server, port = serve(Paraboloid) + self.addCleanup(server.stop, 0) + + client = pmdo.ExplicitClient( + channel=grpc.insecure_channel(f"localhost:{port}") + ) + self.assertIsNone(client.job_id) + + client.run_setup() + + self.assertIsNotNone(client.job_id) + + def test_describe_rpcs_need_no_job(self): + """GetInfo and GetAvailableOptions describe the class, not a run.""" + server, port = serve(Rosenbrock) + self.addCleanup(server.stop, 0) + + client = pmdo.ExplicitClient( + channel=grpc.insecure_channel(f"localhost:{port}") + ) + + client.get_discipline_info() + client.get_available_options() + + self.assertIsNone(client.job_id) + self.assertIn("dimension", client.options_list) + + def test_end_job_releases_the_job(self): + server, port = serve(Paraboloid) + self.addCleanup(server.stop, 0) + + client = pmdo.ExplicitClient( + channel=grpc.insecure_channel(f"localhost:{port}") + ) + client.run_setup() + job_id = client.job_id + + client.end_job() + + self.assertIsNone(client.job_id) + + # the id is genuinely gone from the server + client._job_id = job_id + with self.assertRaises(PhiloteJobError): + client.run_setup() + + def test_job_context_manager_ends_the_job(self): + server, port = serve(Paraboloid) + self.addCleanup(server.stop, 0) + + client = pmdo.ExplicitClient( + channel=grpc.insecure_channel(f"localhost:{port}") + ) + + with client.job(): + self.assertIsNotNone(client.job_id) + + self.assertIsNone(client.job_id) + + def test_keep_alive_defers_eviction(self): + store = JobStore(Discipline, ttl=10.0, sweep_interval=1000.0) + self.addCleanup(store.close_all) + + job = store.create() + job.last_used -= 5.0 + stale = job.last_used + + store.get(job.job_id) + + self.assertGreater(job.last_used, stale) + + def test_set_options_after_setup_is_refused(self): + """ + The metadata was built from the previous values, so a late change + would leave the job describing itself inconsistently. + """ + server, job, _ = make_server(DisciplineServer, Rosenbrock) + context = aborting_job_context(job=job) + + # the legal order: options, then setup + options = data.DisciplineOptions() + options.options.update({"dimension": 3}) + server.SetOptions(options, context) + + server.Setup(data.JobHandle(), context) + self.assertEqual(job.state, JobState.READY) + self.assertEqual( + [tuple(v.shape) for v in job.discipline._var_meta if v.name == "x"], + [(3,)], + ) + + # changing them now would leave the metadata describing the old shape + later = data.DisciplineOptions() + later.options.update({"dimension": 4}) + + with self.assertRaises(Aborted): + server.SetOptions(later, context) + + self.assertEqual( + context.abort.call_args[0][0], grpc.StatusCode.FAILED_PRECONDITION + ) + + +class TestJobErrors(unittest.TestCase): + """The failure paths a client has to be able to tell apart.""" + + def test_missing_header_is_refused(self): + server, job, _ = make_server(DisciplineServer, Paraboloid) + context = aborting_job_context() # no job id at all + + with self.assertRaises(Aborted): + server.Setup(data.JobHandle(), context) + + self.assertEqual( + context.abort.call_args[0][0], grpc.StatusCode.FAILED_PRECONDITION + ) + + def test_unknown_job_is_not_found(self): + server, job, _ = make_server(DisciplineServer, Paraboloid) + context = aborting_job_context(job_id="does-not-exist") + + with self.assertRaises(Aborted): + server.Setup(data.JobHandle(), context) + + self.assertEqual( + context.abort.call_args[0][0], grpc.StatusCode.NOT_FOUND + ) + + def test_client_raises_rather_than_restarting_silently(self): + """ + A client that quietly started a replacement job would run the + optimizer against a fresh discipline and return plausible but wrong + results, so an unknown job has to be terminal. + """ + server, port = serve(Paraboloid) + self.addCleanup(server.stop, 0) + + client = pmdo.ExplicitClient( + channel=grpc.insecure_channel(f"localhost:{port}") + ) + client._job_id = "never-existed" + + with self.assertRaises(PhiloteJobError): + client.run_setup() + + # and it did not paper over the failure by starting a new one + self.assertEqual(client.job_id, "never-existed") + + def test_capacity_is_refused_explicitly(self): + """ + A job can hold a mesh or a solver, so the limit is an explicit refusal + rather than an out-of-memory failure. + """ + server, port = serve(Paraboloid, max_jobs=2) + self.addCleanup(server.stop, 0) + + held = [] + for _ in range(2): + client = pmdo.ExplicitClient( + channel=grpc.insecure_channel(f"localhost:{port}") + ) + client.start_job() + held.append(client) + + third = pmdo.ExplicitClient( + channel=grpc.insecure_channel(f"localhost:{port}") + ) + + with self.assertRaises(Exception) as caught: + third.start_job() + + self.assertIn("maximum", str(caught.exception).lower()) + + # ending one frees the slot + held[0].end_job() + third.start_job() + self.assertIsNotNone(third.job_id) + + +class TestJobStore(unittest.TestCase): + """The store itself, without a server in the way.""" + + def test_rejects_a_non_callable_factory(self): + with self.assertRaises(TypeError): + JobStore(Discipline()) + + def test_each_job_gets_its_own_discipline(self): + store = JobStore(Discipline, ttl=None) + self.addCleanup(store.close_all) + + a, b = store.create(), store.create() + + self.assertIsNot(a.discipline, b.discipline) + self.assertIs(a.discipline.job, a) + self.assertIs(b.discipline.job, b) + + def test_describe_never_hands_out_a_job_instance(self): + store = JobStore(Discipline, ttl=None) + self.addCleanup(store.close_all) + + job = store.create() + + self.assertIsNot(store.describe(), job.discipline) + self.assertIs(store.describe(), store.describe()) + + def test_capacity_error(self): + store = JobStore(Discipline, max_jobs=1, ttl=None) + self.addCleanup(store.close_all) + + store.create() + + with self.assertRaises(JobCapacityError): + store.create() + + def test_unknown_job(self): + store = JobStore(Discipline, ttl=None) + self.addCleanup(store.close_all) + + with self.assertRaises(JobNotFoundError): + store.get("nope") + + def test_sweep_evicts_and_tears_down(self): + torn = [] + + class Tracked(Discipline): + def teardown_job(self): + torn.append(self.job.job_id) + + store = JobStore(Tracked, ttl=0.01, sweep_interval=1000.0) + self.addCleanup(store.close_all) + + job = store.create() + job.last_used -= 1.0 + + self.assertEqual(store.sweep(), [job.job_id]) + self.assertEqual(torn, [job.job_id]) + self.assertEqual(len(store), 0) + + def test_close_runs_teardown(self): + torn = [] + + class Tracked(Discipline): + def teardown_job(self): + torn.append(self.job.job_id) + + store = JobStore(Tracked, ttl=None) + job = store.create() + + store.close(job.job_id) + + self.assertEqual(torn, [job.job_id]) + self.assertEqual(job.state, JobState.CLOSED) + + def test_expiry_during_create_still_runs_teardown(self): + """ + A job reclaimed to make room for a new one must still release what it + held. Reclaiming the slot without tearing down would leak a mesh or a + live solver whenever create() beat the sweeper to an expired job. + """ + torn = [] + + class Tracked(Discipline): + def teardown_job(self): + torn.append(self.job.job_id) + + # no sweeper thread, so create() is the only thing that can reclaim + store = JobStore(Tracked, max_jobs=1, ttl=0.01, sweep_interval=1000.0) + self.addCleanup(store.close_all) + + stale = store.create() + stale.last_used -= 1.0 + + fresh = store.create() + + self.assertEqual(torn, [stale.job_id]) + self.assertEqual(len(store), 1) + self.assertNotEqual(fresh.job_id, stale.job_id) + + def test_failed_construction_frees_the_slot(self): + calls = [] + + def factory(): + calls.append(1) + raise RuntimeError("mesh missing") + + store = JobStore(factory, max_jobs=1, ttl=None) + self.addCleanup(store.close_all) + + for _ in range(2): + with self.assertRaises(RuntimeError): + store.create() + + # the first failure did not permanently consume the only slot + self.assertEqual(len(calls), 2) + self.assertEqual(len(store), 0) + + +class TestThreadPoolWarning(unittest.TestCase): + """max_jobs is not the cap that binds; the gRPC thread pool is.""" + + def test_warns_when_pool_is_smaller_than_job_cap(self): + grpc_server = grpc.server(futures.ThreadPoolExecutor(max_workers=2)) + self.addCleanup(grpc_server.stop, 0) + + server = pmdo.ExplicitServer(discipline=Paraboloid, max_jobs=8) + + with self.assertWarns(RuntimeWarning) as caught: + server.attach_to_server(grpc_server) + + self.assertIn("thread pool", str(caught.warning)) + + def test_quiet_when_pool_is_large_enough(self): + import warnings + + grpc_server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) + self.addCleanup(grpc_server.stop, 0) + + server = pmdo.ExplicitServer(discipline=Paraboloid, max_jobs=8) + + with warnings.catch_warnings(): + warnings.simplefilter("error") + server.attach_to_server(grpc_server) + + +class FakeRpcError(grpc.RpcError): + """A gRPC error with a chosen status code, for the client error paths.""" + + def __init__(self, code, details="boom"): + self._code = code + self._details = details + + def code(self): + return self._code + + def details(self): + return self._details + + +class TestClientErrorTranslation(unittest.TestCase): + """ + Every client call must surface a server failure as a Philote exception. + + Before this, only the compute calls translated, so a failure during the + setup phase escaped as a raw grpc.RpcError and callers could not tell an + expired job from a genuine server fault. + """ + + def _client(self): + client = pmdo.ExplicitClient(channel=Mock()) + client._job_id = "job-1" + client._disc_stub = Mock() + return client + + def test_not_found_becomes_job_not_found(self): + for method, call in ( + ("GetInfo", lambda c: c.get_discipline_info()), + ("SetStreamOptions", lambda c: c.send_stream_options()), + ("GetAvailableOptions", lambda c: c.get_available_options()), + ("SetOptions", lambda c: c.send_options({"a": 1})), + ("Setup", lambda c: c.run_setup()), + ("GetVariableDefinitions", lambda c: c.get_variable_definitions()), + ("GetPartialDefinitions", lambda c: c.get_partials_definitions()), + ("SetVariableShapes", lambda c: c.send_variable_shapes([])), + ("EndJob", lambda c: c.end_job()), + ("KeepAlive", lambda c: c.keep_alive()), + ): + with self.subTest(rpc=method): + client = self._client() + getattr(client._disc_stub, method).side_effect = FakeRpcError( + grpc.StatusCode.NOT_FOUND + ) + + with self.assertRaises(JobNotFoundError): + call(client) + + def test_resource_exhausted_becomes_capacity_error(self): + client = self._client() + client._job_id = None + client._disc_stub.StartJob.side_effect = FakeRpcError( + grpc.StatusCode.RESOURCE_EXHAUSTED + ) + + with self.assertRaises(JobCapacityError): + client.start_job() + + def test_other_codes_stay_server_errors(self): + client = self._client() + client._disc_stub.Setup.side_effect = FakeRpcError( + grpc.StatusCode.INTERNAL + ) + + with self.assertRaises(PhiloteServerError): + client.run_setup() + + def test_start_job_is_idempotent(self): + client = self._client() + + self.assertEqual(client.start_job(), "job-1") + client._disc_stub.StartJob.assert_not_called() + + def test_end_job_without_a_job_is_a_no_op(self): + client = self._client() + client._job_id = None + + client.end_job() + + client._disc_stub.EndJob.assert_not_called() + + def test_keep_alive_calls_the_rpc(self): + client = self._client() + + client.keep_alive() + + client._disc_stub.KeepAlive.assert_called_once() + + +class TestServerWithoutAFactory(unittest.TestCase): + """ + A server built with no factory has to say so, not fail obscurely. + """ + + def test_every_entry_point_refuses(self): + server = pmdo.ExplicitServer() + + for name, call in ( + ("StartJob", lambda c: server.StartJob(data.JobHandle(), c)), + ("Setup", lambda c: server.Setup(data.JobHandle(), c)), + ("GetInfo", lambda c: server.GetInfo(data.JobHandle(), c)), + ("GetAvailableOptions", + lambda c: server.GetAvailableOptions(data.JobHandle(), c)), + ): + with self.subTest(rpc=name): + context = aborting_job_context(job_id="whatever") + + with self.assertRaises(Aborted): + call(context) + + self.assertEqual( + context.abort.call_args[0][0], + grpc.StatusCode.FAILED_PRECONDITION, + ) + + +class TestServerJobRpcs(unittest.TestCase): + """StartJob, EndJob and KeepAlive as handlers.""" + + def test_end_job_twice_is_not_found(self): + server, job, context = make_server(DisciplineServer, Paraboloid) + + server.EndJob(data.JobHandle(), context) + + second = aborting_job_context(job_id=job.job_id) + with self.assertRaises(Aborted): + server.EndJob(data.JobHandle(), second) + + self.assertEqual( + second.abort.call_args[0][0], grpc.StatusCode.NOT_FOUND + ) + + def test_end_job_handles_a_concurrent_close(self): + """ + EndJob resolves the job and then closes it, so another thread can + close it in between. The handler for that race is not dead code. + """ + server, job, context = make_server(DisciplineServer, Paraboloid) + aborting = aborting_job_context(job=job) + server._jobs.close = Mock(side_effect=JobNotFoundError("raced")) + + with self.assertRaises(Aborted): + server.EndJob(data.JobHandle(), aborting) + + self.assertEqual( + aborting.abort.call_args[0][0], grpc.StatusCode.NOT_FOUND + ) + + def test_end_job_reports_an_unexpected_teardown_failure(self): + """A discipline whose teardown_job raises must not be silent.""" + server, job, context = make_server(DisciplineServer, Paraboloid) + aborting = aborting_job_context(job=job) + server._jobs.close = Mock(side_effect=RuntimeError("solver stuck")) + + with self.assertRaises(Aborted): + server.EndJob(data.JobHandle(), aborting) + + self.assertEqual( + aborting.abort.call_args[0][0], grpc.StatusCode.INTERNAL + ) + self.assertIn("EndJob failed", aborting.abort.call_args[0][1]) + + def test_keep_alive_refreshes_the_job(self): + server, job, context = make_server(DisciplineServer, Paraboloid) + job.last_used -= 5.0 + stale = job.last_used + + server.KeepAlive(data.JobHandle(), context) + + self.assertGreater(job.last_used, stale) + + def test_start_job_reports_capacity(self): + server, job, _ = make_server(DisciplineServer, Paraboloid, max_jobs=1) + context = aborting_job_context(job=job) + + with self.assertRaises(Aborted): + server.StartJob(data.JobHandle(), context) + + self.assertEqual( + context.abort.call_args[0][0], grpc.StatusCode.RESOURCE_EXHAUSTED + ) + + def test_start_job_reports_a_failing_factory(self): + def broken(): + raise RuntimeError("mesh missing") + + server = pmdo.ExplicitServer(discipline=broken, ttl=None) + context = aborting_job_context(job_id="x") + + with self.assertRaises(Aborted): + server.StartJob(data.JobHandle(), context) + + self.assertEqual( + context.abort.call_args[0][0], grpc.StatusCode.INTERNAL + ) + + def test_unmapped_job_error_falls_back_to_internal(self): + from philote_mdo.general.discipline_server import _job_status + + self.assertEqual( + _job_status(PhiloteJobError("unclassified")), + grpc.StatusCode.INTERNAL, + ) + + +class TestJobStoreDetails(unittest.TestCase): + """Remaining JobStore behaviour.""" + + def test_max_jobs_is_reported(self): + store = JobStore(Discipline, max_jobs=3, ttl=None) + self.addCleanup(store.close_all) + + self.assertEqual(store.max_jobs, 3) + + def test_passing_an_instance_says_how_to_fix_it(self): + """ + `discipline=` took an instance before this change and takes a class + now, so the same call means something different. An instance is not + callable, so it fails immediately rather than silently -- and the + message has to say what to write instead. + """ + with self.assertRaises(TypeError) as caught: + pmdo.ExplicitServer(discipline=Paraboloid()) + + message = str(caught.exception) + self.assertIn("discipline=Paraboloid)", message) + self.assertIn("not", message) + self.assertIn("discipline=Paraboloid())", message) + + def test_closing_an_unknown_job_raises(self): + store = JobStore(Discipline, ttl=None) + self.addCleanup(store.close_all) + + with self.assertRaises(JobNotFoundError): + store.close("nope") + + def test_teardown_of_a_half_built_job_is_safe(self): + """A job whose factory never finished has no discipline to release.""" + store = JobStore(Discipline, ttl=None) + self.addCleanup(store.close_all) + + job = store.create() + job.discipline = None + + store.close(job.job_id) + + self.assertEqual(job.state, JobState.CLOSED) + + def test_sweeper_thread_evicts_without_being_asked(self): + torn = [] + + class Tracked(Discipline): + def teardown_job(self): + torn.append(self.job.job_id) + + store = JobStore(Tracked, ttl=0.05, sweep_interval=0.02) + self.addCleanup(store.close_all) + + job = store.create() + + deadline = time.monotonic() + 5.0 + while len(store) and time.monotonic() < deadline: + time.sleep(0.02) + + self.assertEqual(len(store), 0) + self.assertEqual(torn, [job.job_id]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_openmdao_explicit_client.py b/tests/test_openmdao_explicit_client.py index daa6a0e..abd63d2 100644 --- a/tests/test_openmdao_explicit_client.py +++ b/tests/test_openmdao_explicit_client.py @@ -28,6 +28,8 @@ # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. import unittest + +from conftest import patch_discipline_stub from unittest.mock import Mock, MagicMock, patch import philote_mdo.generated.data_pb2 as data from philote_mdo.openmdao import RemoteExplicitComponent @@ -39,6 +41,16 @@ class TestOpenMdaoExplicitClient(unittest.TestCase): Unit tests for the OpenMDAO explicit component/client. """ + def setUp(self): + # the component claims a job from inside __init__, so the stub has + # to be mocked before construction rather than after + self._job_patch = patch_discipline_stub() + self._job_patch.start() + + def tearDown(self): + self._job_patch.stop() + + @patch("philote_mdo.general.ExplicitClient") def test_constructor(self, mock_explicit_client, mock_explicit_component): """ diff --git a/tests/test_openmdao_implicit_client.py b/tests/test_openmdao_implicit_client.py index cbc2246..2ca5105 100644 --- a/tests/test_openmdao_implicit_client.py +++ b/tests/test_openmdao_implicit_client.py @@ -28,6 +28,8 @@ # therein. The DoD does not exercise any editorial, security, or other # control over the information you may find at these locations. import unittest + +from conftest import patch_discipline_stub from unittest.mock import Mock, MagicMock, patch import numpy as np import philote_mdo.generated.data_pb2 as data @@ -40,6 +42,16 @@ class TestOpenMdaoImplicitClient(unittest.TestCase): Unit tests for the OpenMDAO implicit component/client. """ + def setUp(self): + # the component claims a job from inside __init__, so the stub has + # to be mocked before construction rather than after + self._job_patch = patch_discipline_stub() + self._job_patch.start() + + def tearDown(self): + self._job_patch.stop() + + @patch("philote_mdo.general.ImplicitClient") def test_constructor(self, mock_explicit_client, mock_implicit_component): """ diff --git a/tests/test_openmdao_integration.py b/tests/test_openmdao_integration.py index 677b09c..a932ed5 100644 --- a/tests/test_openmdao_integration.py +++ b/tests/test_openmdao_integration.py @@ -29,6 +29,8 @@ # control over the information you may find at these locations. from concurrent import futures import unittest + +from conftest import job_context, make_server import grpc import numpy as np from numpy.testing import assert_almost_equal @@ -48,9 +50,9 @@ def test_openmdao_paraboloid_compute(self): Integration test for the Paraboloid compute function. """ # server code - server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) - discipline = pmdo.ExplicitServer(discipline=Paraboloid()) + discipline = pmdo.ExplicitServer(discipline=Paraboloid) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -83,9 +85,9 @@ def test_paraboloid_compute_partials(self): Integration test for the Paraboloid compute function. """ # server code - server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) - discipline = pmdo.ExplicitServer(discipline=Paraboloid()) + discipline = pmdo.ExplicitServer(discipline=Paraboloid) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -119,9 +121,9 @@ def test_rosenbrock_compute(self): Integration test for the Paraboloid compute function. """ # server code - server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) - discipline = pmdo.ExplicitServer(discipline=Rosenbrock()) + discipline = pmdo.ExplicitServer(discipline=Rosenbrock) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -157,9 +159,9 @@ def test_rosenbrock_option_set_after_construction(self): shows up as a shape mismatch (issue #77). """ # server code - server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) - discipline = pmdo.ExplicitServer(discipline=Rosenbrock()) + discipline = pmdo.ExplicitServer(discipline=Rosenbrock) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -196,9 +198,9 @@ def test_rosenbrock_compute_partials(self): Integration test for the Paraboloid compute function. """ # server code - server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) - discipline = pmdo.ExplicitServer(discipline=Rosenbrock()) + discipline = pmdo.ExplicitServer(discipline=Rosenbrock) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051") @@ -233,9 +235,9 @@ def test_quadratic_compute_function(self): example. """ # server code - server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) + server = grpc.server(futures.ThreadPoolExecutor(max_workers=16)) - discipline = pmdo.ImplicitServer(discipline=QuadradicImplicit()) + discipline = pmdo.ImplicitServer(discipline=QuadradicImplicit) discipline.attach_to_server(server) server.add_insecure_port("[::]:50051")