rlmesh 0.1.0

Internal RLMesh crate (unstable Rust API): Rust bindings for model-environment evaluation; build on the rlmesh Python package.
Documentation
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
use std::sync::Arc;
use std::time::{Duration, Instant};

use async_trait::async_trait;
use rlmesh_grpc::wire::{
    encode_batched_partial_values, env_contract_from_proto, env_contract_to_proto,
};
use rlmesh_proto::model::v1::{PredictRequest, ResetAdapterRequest};
use rlmesh_proto::{EndpointPhases, elapsed_ns};
use rlmesh_runtime::{
    PeerCeiling, RuntimeDriver, RuntimeEnv, RuntimeEnvReset, RuntimeEnvStep, RuntimeError,
    RuntimeHooks, RuntimeModel, RuntimeModelPrediction, RuntimeReport, RuntimeSessionSpec,
};

use super::handler::{ModelHandler, PredictFrames};
use super::wire::{
    ModelAction, check_actions_conform, encode_replay_frames, model_action_to_endpoint_response,
    model_observation_from_endpoint_request,
};
use crate::{Error, Result, spaces};

/// Connect to the env, resolve the handler's route, and drive the runtime
/// loop to completion. The route resolves before driving, exactly as the
/// served path does at `ResolveAdapter`: that arms the spec'd engine path
/// (adapter, per-episode frame buffers, batched/chunked predict corners) and
/// pins the execution horizon — without it every predict would fall into the
/// spec-less branch. The handler is dropped when the run ends, so there is no
/// release step.
pub(super) async fn run_local<H>(
    handler: &mut H,
    options: crate::RunLocalOptions,
    cancellation: tokio_util::sync::CancellationToken,
    hooks: Arc<dyn RuntimeHooks>,
) -> Result<RuntimeReport>
where
    H: ModelHandler + 'static,
{
    let mut env = rlmesh_grpc::EnvClient::connect_with_token(
        &options.env_address.to_string(),
        &options.token,
    )
    .await
    .map_err(Error::from)?;
    // The runtime tier's WANT rides this leg's handshake and caps the env-leg
    // negotiation, so it has to be declared before the handshake goes out.
    env.declare_workflow_edition(options.workflow_edition.clone());
    let handshake = env.handshake().await.map_err(Error::from)?;
    let env_ceiling = env_ceiling(&handshake);
    // A lane endpoint steps/resets lanes individually, so a num_envs > 1
    // session can run driver-owned (DISABLED) resets lane by lane.
    let subset_step = rlmesh_proto::has_capability(
        &handshake.capabilities,
        rlmesh_proto::capabilities::ENV_SUBSET_STEP,
    );
    let env_contract = env_contract_from_proto(handshake.env_contract)
        .map_err(|err| Error::Internal(format!("invalid spaces spec from env: {err}")))?;
    // The handler returns typed actions the runtime encodes against this space;
    // a missing one would fail every predict, so reject at connect, not mid-run.
    if env_contract.action_space.is_none() {
        return Err(Error::Internal(
            "env contract has no action_space; a model cannot encode actions without it"
                .to_string(),
        ));
    }
    // The runtime is the env-id authority (R1); for the in-process path mint a
    // UUIDv7 container id. The human env name lives on the SDK's own contract.
    let env_id = crate::mint_id();
    let num_envs = handshake.num_envs;
    let session_id = format!("local-{}", std::process::id());

    // Action chunking across a lockstep vector env would replay one whole-batch
    // chunk for every lane, and a lane that ends mid-chunk invalidates the buffer
    // for all of them. A lane endpoint replays per lane, so it is fine there.
    if num_envs > 1 && !subset_step && options.execution_horizon > 1 {
        return Err(Error::Internal(format!(
            "execution_horizon={} cannot be combined with a lockstep vector env (num_envs={num_envs}): \
             chunk replay is whole-batch, so one lane's episode end discards every lane's \
             buffered frames. Use num_envs=1, a lane endpoint, or execution_horizon=1.",
            options.execution_horizon,
        )));
    }

    // The driver delivers every replayed step as a history row, so a stacked
    // adapter's window advances on every step at any horizon; the route answers
    // whether it needs that, and the driver only buffers rows when it does. The
    // handler is in-process (this build's engine), so unlike the served path
    // there is no peer capability to gate the offer on.
    let mut wants_history = false;
    if let Some(route_setup) = handler.route_setup() {
        let needs = route_setup
            .resolve_adapter(
                &env_id,
                &env_contract,
                crate::model::ResolveOptions {
                    execution_horizon: options.execution_horizon,
                    delivers_history: true,
                },
            )
            .await?;
        wants_history = needs.history.is_some();
    }

    let spec = RuntimeSessionSpec {
        session_id,
        env_id,
        env_component_id: "local-env".to_string(),
        model_component_id: "local-model".to_string(),
        workflow_edition: handshake.workflow_edition,
        env_contract: env_contract_to_proto(&env_contract),
        num_envs,
        base_seed: options.base_seed,
        episode_seeds: options.episode_seeds,
        max_episodes: options.max_episodes,
        trial_index_base: options.trial_index_base,
        max_episode_steps: options.max_episode_steps,
        max_episode_seconds: options.max_episode_seconds,
        close_env_on_end: options.close_env,
        // A lane endpoint is driven one group per lane; the grouped predicts
        // reach the handler's `predict_grouped` (one fused forward for a model
        // with a batched corner).
        subset_step,
        limits: Default::default(),
        env_ceiling: Some(env_ceiling),
        model_ceiling: None,
    };
    let env = EnvClientRuntimeEnv::new(env);
    let model = ModelHandlerRuntimeModel::new(handler, env_contract).with_history(wants_history);
    RuntimeDriver::new(spec, env, model, hooks)
        .with_prefetch(options.prefetch_lead)
        .run_with_cancellation_reason(cancellation, "interrupted by the host (signal)")
        .await
        .map_err(run_error)
}

/// What the served env on the other end of `handshake` can decode.
pub(super) fn env_ceiling(handshake: &rlmesh_grpc::EnvHandshake) -> PeerCeiling {
    PeerCeiling::wire_v1(
        PeerCeiling::highest_shared_edition(&handshake.supported_workflow_editions)
            .unwrap_or(handshake.workflow_edition),
        handshake.capabilities.clone(),
        rlmesh_grpc::MAX_MESSAGE_SIZE,
    )
}

/// Wrap a facade [`Error`] as a driver model-RPC failure, keeping its
/// recoverable flag (the default `model_rpc` constructor drops it, so a
/// recoverable handler decline used to reach the caller as permanent).
fn model_rpc(error: Error) -> RuntimeError {
    RuntimeError::model_rpc_with_recoverability("local-model", error.is_recoverable(), error)
}

/// Map the driver's error back onto the facade taxonomy. Flattening every
/// failure to [`Error::Internal`] made a model decline, an env fault, a dropped
/// connection and a timeout indistinguishable to the caller (and always
/// non-recoverable); the structured `#[source]` each RPC variant carries is the
/// original error, so unwrap it when it is one of ours.
fn run_error(error: RuntimeError) -> Error {
    let recoverable = error.is_recoverable();
    let message = error.to_string();
    match error {
        RuntimeError::ModelRpc { source, .. } => source
            .and_then(|source| source.downcast::<Error>().ok())
            .map_or_else(
                || {
                    if recoverable {
                        Error::model_recoverable(message)
                    } else {
                        Error::model(message)
                    }
                },
                |error| *error,
            ),
        // Keep the driver's message (it names the op and the step) and take only
        // the classification from the structured source.
        RuntimeError::EnvRpc { source, .. } => match source
            .and_then(|source| source.downcast::<rlmesh_grpc::error::Error>().ok())
            .map(|error| Error::from(*error))
        {
            Some(Error::Connection(_)) => Error::Connection(message),
            Some(Error::Timeout(timeout)) => Error::Timeout(timeout),
            Some(Error::Environment(env)) => {
                Error::Environment(crate::EnvironmentError { message, ..env })
            }
            _ => Error::Environment(crate::EnvironmentError {
                code: crate::ErrorCode::Internal,
                message,
                is_recoverable: recoverable,
            }),
        },
        RuntimeError::OperationTimeout { timeout, .. } => Error::Timeout(timeout),
        _ => Error::Internal(message),
    }
}

/// Adapts a connected [`rlmesh_grpc::EnvClient`] to the [`RuntimeEnv`] trait
/// expected by [`rlmesh_runtime::RuntimeDriver`].
///
/// Use this to drive a remote environment from your own `RuntimeDriver`
/// embedding without re-implementing the per-call telemetry choreography: the
/// adapter takes the client's last-operation telemetry after each `reset`/`step`
/// and attaches it to the runtime result for you, and it maps transport errors
/// onto recoverable/non-recoverable [`rlmesh_runtime::RuntimeError`]s.
#[derive(Clone)]
pub struct EnvClientRuntimeEnv {
    inner: rlmesh_grpc::EnvClient,
}

impl EnvClientRuntimeEnv {
    /// Wrap a connected (and handshaked) env client.
    pub fn new(client: rlmesh_grpc::EnvClient) -> Self {
        Self { inner: client }
    }

    /// Consume the adapter and return the underlying client.
    pub fn into_inner(self) -> rlmesh_grpc::EnvClient {
        self.inner
    }
}

#[async_trait]
impl RuntimeEnv for EnvClientRuntimeEnv {
    async fn reset(
        &mut self,
        request: rlmesh_proto::env::v1::ResetRequest,
    ) -> std::result::Result<RuntimeEnvReset, rlmesh_runtime::RuntimeError> {
        let response = self.inner.reset(request).await.map_err(|err| {
            let recoverable = err.is_recoverable();
            rlmesh_runtime::RuntimeError::env_rpc_with_recoverability(
                "env.reset",
                0,
                recoverable,
                err,
            )
        })?;
        Ok(RuntimeEnvReset {
            response,
            endpoint_total_ns: self.inner.take_last_endpoint_total_ns(),
            phases: self.inner.take_last_phases(),
        })
    }

    async fn step(
        &mut self,
        request: rlmesh_proto::env::v1::StepRequest,
    ) -> std::result::Result<RuntimeEnvStep, rlmesh_runtime::RuntimeError> {
        let response = self.inner.step(request).await.map_err(|err| {
            let recoverable = err.is_recoverable();
            rlmesh_runtime::RuntimeError::env_rpc_with_recoverability(
                "env.step",
                0,
                recoverable,
                err,
            )
        })?;
        Ok(RuntimeEnvStep {
            response,
            endpoint_total_ns: self.inner.take_last_endpoint_total_ns(),
            phases: self.inner.take_last_phases(),
        })
    }

    async fn close(&mut self, timeout: Duration) -> std::result::Result<(), String> {
        let close = self.inner.close();
        tokio::time::timeout(timeout, close)
            .await
            .map_err(|err| err.to_string())?
            .map(|_| ())
            .map_err(|err| err.to_string())
    }
}

/// Adapts a [`ModelHandler`] to the [`RuntimeModel`] trait expected by
/// [`rlmesh_runtime::RuntimeDriver`].
///
/// Use this to drive your handler from your own `RuntimeDriver` embedding. The
/// adapter decodes the runtime's predict request into a
/// [`ModelObservation`](crate::ModelObservation), runs `predict`, and re-encodes
/// the action, matching the choreography the in-process `run_local` path
/// performs. Per-episode lifecycle is explicit (see below) — there is no
/// episode-begin hook; the model's state is lazy-seeded on first predict.
///
/// It borrows the handler mutably so the caller retains ownership (e.g. to run
/// the close hook afterward). Per-episode lifecycle is explicit (R2): the runtime
/// driver emits `ResetAdapter` on episode end, routed here to the handler's
/// `reset_adapter`; there is no position-diff / active-episodes state.
pub struct ModelHandlerRuntimeModel<'a, H> {
    /// The handler, behind a lock: the driver may evict adapter state while a
    /// grouped predict is being prepared, and a served handler is shared the
    /// same way.
    handler: Arc<tokio::sync::Mutex<&'a mut H>>,
    env_contract: Arc<spaces::EnvContract>,
    /// The route asked for observation history at resolve (see
    /// [`with_history`](Self::with_history)).
    wants_history: bool,
}

impl<'a, H> ModelHandlerRuntimeModel<'a, H> {
    /// Build an adapter for `handler` against the given env contract.
    pub fn new(handler: &'a mut H, env_contract: spaces::EnvContract) -> Self {
        Self {
            handler: Arc::new(tokio::sync::Mutex::new(handler)),
            env_contract: Arc::new(env_contract),
            wants_history: false,
        }
    }

    /// Tell the driver whether the route negotiated observation history
    /// (`RouteNeeds::history` was answered at resolve), so it buffers every
    /// replayed step as a history row on the next predict.
    pub fn with_history(mut self, wants_history: bool) -> Self {
        self.wants_history = wants_history;
        self
    }
}

/// One group of a grouped predict, prepared for the handler: what encoding
/// its reply needs.
struct PreparedGroup {
    route: crate::model::types::ModelRouteContext,
    num_envs: usize,
}

#[async_trait]
impl<H> RuntimeModel for ModelHandlerRuntimeModel<'_, H>
where
    H: ModelHandler + 'static,
{
    fn wants_history(&self) -> bool {
        self.wants_history
    }

    /// `predict_group` below is one fused forward, so the driver batches the
    /// lanes waiting on it rather than predicting each on its own.
    fn fuses_predicts(&self) -> bool {
        true
    }

    async fn predict(
        &self,
        request: PredictRequest,
    ) -> std::result::Result<RuntimeModelPrediction, rlmesh_runtime::RuntimeError> {
        self.predict_group(vec![request])
            .await
            .pop()
            .unwrap_or_else(|| {
                Err(rlmesh_runtime::RuntimeError::model_rpc(
                    "local-model",
                    Error::model("predict_grouped returned no result"),
                ))
            })
    }

    /// Every request decodes into one `ModelObservation` and the batch goes to
    /// the handler's `predict_grouped` — the same seam the served endpoint
    /// uses, so a model with a batched corner runs one fused forward for all
    /// the lanes waiting on it. Results align 1:1 and in order.
    async fn predict_group(
        &self,
        requests: Vec<PredictRequest>,
    ) -> Vec<std::result::Result<RuntimeModelPrediction, rlmesh_runtime::RuntimeError>> {
        let started = Instant::now();
        let model_err = model_rpc;
        let Some(action_space) = self.env_contract.action_space.clone() else {
            return requests
                .iter()
                .map(|_| {
                    Err(model_err(Error::model(
                        "model route contract missing action space",
                    )))
                })
                .collect();
        };
        // Prepare every group; one that fails to decode reports its own error
        // and is left out of the batch.
        let mut batch = Vec::with_capacity(requests.len());
        let mut prepared: Vec<std::result::Result<PreparedGroup, rlmesh_runtime::RuntimeError>> =
            Vec::with_capacity(requests.len());
        for request in requests {
            match model_observation_from_endpoint_request(request) {
                Ok(mut observation) => {
                    // The request's row count is its own width (a lane group
                    // sends one row), never the route's.
                    let num_envs = observation.route.episodes.len().max(1);
                    observation.env_contract = Some(Arc::clone(&self.env_contract));
                    observation.num_envs = num_envs;
                    prepared.push(Ok(PreparedGroup {
                        route: observation.route.clone(),
                        num_envs,
                    }));
                    batch.push(observation);
                }
                Err(err) => prepared.push(Err(model_err(err))),
            }
        }
        let decode_ns = elapsed_ns(started);

        let mut handler = self.handler.lock().await;
        let call_started = Instant::now();
        let mut frames = handler.predict_grouped(batch).await.into_iter();
        let user_ns = elapsed_ns(call_started);
        // Drain the adapter share even for a failed forward, or its time leaks
        // into the next predict's `adapter_ns`.
        let adapter_ns = handler.take_adapter_ns();
        let held = handler.held_state();
        drop(handler);

        let encode_started = Instant::now();
        prepared
            .into_iter()
            .map(|group| {
                let PreparedGroup { route, num_envs } = group?;
                let PredictFrames { actions, replay } = frames
                    .next()
                    .ok_or_else(|| {
                        model_err(Error::model(
                            "predict_grouped returned fewer results than prepared groups",
                        ))
                    })?
                    .map_err(model_err)?;
                if actions.len() != num_envs {
                    return Err(model_err(Error::model(format!(
                        "predict returned {} actions for {num_envs} lanes",
                        actions.len()
                    ))));
                }
                check_actions_conform(&action_space, &actions).map_err(model_err)?;
                let frame0 = encode_batched_partial_values(&actions, &action_space)
                    .map_err(|err| model_err(Error::model(err.to_string())))?;
                let replay_frames =
                    encode_replay_frames(&replay, num_envs, &action_space).map_err(model_err)?;
                let mut wire_actions = Vec::with_capacity(1 + replay_frames.len());
                wire_actions.push(frame0);
                wire_actions.extend(replay_frames);
                Ok(RuntimeModelPrediction {
                    response: model_action_to_endpoint_response(ModelAction {
                        actions: wire_actions,
                        route,
                    }),
                    endpoint_total_ns: Some(elapsed_ns(started)),
                    phases: EndpointPhases {
                        decode_ns,
                        user_ns,
                        encode_ns: elapsed_ns(encode_started),
                        adapter_ns,
                        held_episodes: held
                            .map(|held| held.episodes.min(u64::from(u32::MAX)) as u32),
                        held_state_bytes: held.map(|held| held.bytes),
                        ..EndpointPhases::default()
                    },
                    group_size: None,
                })
            })
            .collect()
    }

    async fn reset_adapter(
        &self,
        request: ResetAdapterRequest,
    ) -> std::result::Result<(), RuntimeError> {
        // Route the driver's explicit episode-end GC to the route setup's
        // eviction and then the handler's own hook, as the served path does.
        let env_id = request
            .context
            .map(|context| context.env_id)
            .unwrap_or_default();
        let mut handler = self.handler.lock().await;
        if let Some(route_setup) = handler.route_setup() {
            route_setup
                .reset_adapter(&env_id, &request.episode_ids)
                .await
                .map_err(model_rpc)?;
        }
        handler
            .reset_adapter(&env_id, request.episode_ids)
            .await
            .map_err(model_rpc)
    }
}