Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
70 changes: 61 additions & 9 deletions codex-rs/core/src/context/world_state/environment.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@ use codex_protocol::protocol::TurnContextItem;
use codex_protocol::protocol::TurnContextNetworkItem;
use codex_utils_absolute_path::AbsolutePathBuf;
use codex_utils_path_uri::PathUri;
use serde::Deserialize;
use serde::Serialize;
use std::collections::BTreeMap;

/// Environment values visible to the model.
Expand Down Expand Up @@ -91,17 +93,49 @@ impl EnvironmentsState {
}

impl WorldStateSection for EnvironmentsState {
fn render_diff(&self, previous: Option<&Self>) -> Option<Box<dyn ContextualUserFragment>> {
let empty = Self::default();
const ID: &'static str = "environments";
type Snapshot = EnvironmentsSnapshot;

fn snapshot(&self) -> Self::Snapshot {
EnvironmentsSnapshot {
environments: self
.environments
.iter()
.map(|(id, environment)| {
(
id.clone(),
EnvironmentSnapshot {
cwd: environment.cwd.inferred_native_path_string(),
status: environment.status,
shell: environment.shell.clone(),
},
)
})
.collect(),
current_date: self.current_date.clone(),
timezone: self.timezone.clone(),
network: self.network.as_ref().map(NetworkContext::render),
filesystem: self.filesystem.as_ref().map(FileSystemContext::render),
subagents: self.subagents.clone(),
}
}

fn render_diff(
&self,
previous: Option<&Self::Snapshot>,
) -> Option<Box<dyn ContextualUserFragment>> {
let current = self.snapshot();
let empty = EnvironmentsSnapshot::default();
let previous = previous.unwrap_or(&empty);
let turn_context_values_changed = self.current_date != previous.current_date
|| self.timezone != previous.timezone
|| self.network != previous.network
|| self.filesystem != previous.filesystem;
let turn_context_values_changed = current.current_date != previous.current_date
|| current.timezone != previous.timezone
|| current.network != previous.network
|| current.filesystem != previous.filesystem;
let mut updates = self
.environments
.iter()
.filter(|(id, environment)| {
.filter(|(id, _)| {
let environment = &current.environments[*id];
previous
.environments
.get(*id)
Expand Down Expand Up @@ -269,7 +303,24 @@ struct EnvironmentState {
shell: Option<String>,
}

impl EnvironmentState {
#[derive(Default, Deserialize, Serialize)]
pub(crate) struct EnvironmentsSnapshot {
environments: BTreeMap<String, EnvironmentSnapshot>,
current_date: Option<String>,
timezone: Option<String>,
network: Option<String>,
filesystem: Option<String>,
subagents: Option<String>,
}

#[derive(Deserialize, Serialize)]
struct EnvironmentSnapshot {
cwd: String,
status: EnvironmentStatus,
shell: Option<String>,
}

impl EnvironmentSnapshot {
fn has_same_diff_value(&self, other: &Self) -> bool {
self.cwd == other.cwd
&& self.status == other.status
Expand All @@ -281,7 +332,8 @@ impl EnvironmentState {
}
}

#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
enum EnvironmentStatus {
Starting,
Available,
Expand Down
42 changes: 40 additions & 2 deletions codex-rs/core/src/context/world_state/environment_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ use codex_protocol::models::PermissionProfile;
use codex_protocol::models::ResponseItem;
use codex_protocol::permissions::NetworkSandboxPolicy;
use pretty_assertions::assert_eq;
use serde_json::json;

#[test]
fn renders_full_environment_state() -> Result<()> {
Expand Down Expand Up @@ -95,7 +96,7 @@ fn renders_only_changed_environments() -> Result<()> {
</environments>
</environment_context>"#,
)],
render_fragments(current.render_diff(&previous)),
render_fragments(current.render_diff(&previous.snapshot())),
);
Ok(())
}
Expand Down Expand Up @@ -151,7 +152,42 @@ fn persisted_turn_context_values_render_a_diff() -> Result<()> {
<filesystem><permission_profile type="external"><file_system type="external" /></permission_profile></filesystem>
</environment_context>"#,
)],
render_fragments(current.render_diff(&previous)),
render_fragments(current.render_diff(&previous.snapshot())),
);
Ok(())
}

#[test]
fn persisted_snapshot_uses_model_visible_path_and_context_values() -> Result<()> {
let mut world_state = WorldState::default();
world_state.add_section(EnvironmentsState {
environments: [(
"remote".to_string(),
available("file:///C:/windows", "powershell")?,
)]
.into_iter()
.collect(),
filesystem: Some(FileSystemContext::from_permission_profile(
&PermissionProfile::Disabled,
&[],
)),
..Default::default()
});

assert_eq!(
serde_json::to_value(world_state.snapshot())?,
json!({
"environments": {
"environments": {
"remote": {
"cwd": "C:\\windows",
"status": "available",
"shell": "powershell"
}
},
"filesystem": "<filesystem><permission_profile type=\"disabled\"><file_system type=\"unrestricted\" /></permission_profile></filesystem>"
}
}),
);
Ok(())
}
Expand Down Expand Up @@ -180,6 +216,7 @@ fn single_environment_diff_ignores_unknown_shell() -> Result<()> {
.collect(),
..Default::default()
};
let previous = WorldStateSection::snapshot(&previous);

assert_eq!(
None,
Expand All @@ -199,6 +236,7 @@ fn removed_legacy_environment_renders_unavailable() -> Result<()> {
.collect(),
..Default::default()
};
let previous = WorldStateSection::snapshot(&previous);

assert_eq!(
Some(user_message(
Expand Down
130 changes: 97 additions & 33 deletions codex-rs/core/src/context/world_state/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,50 +2,82 @@ mod environment;

use crate::context::ContextualUserFragment;
use indexmap::IndexMap;
use std::any::Any;
use std::any::TypeId;
use serde::Serialize;
use serde::de::DeserializeOwned;
use serde_json::Value;
use std::collections::BTreeMap;
use std::fmt;

pub(crate) use environment::EnvironmentsState;

trait ErasedWorldStateSection: Send + Sync {
fn as_any(&self) -> &dyn Any;
fn snapshot(&self) -> Option<Value>;

fn render_diff(&self, previous: Option<&dyn Any>) -> Option<Box<dyn ContextualUserFragment>>;
fn render_diff(&self, previous: Option<&Value>) -> Option<Box<dyn ContextualUserFragment>>;
}

impl<S: WorldStateSection> ErasedWorldStateSection for S {
fn as_any(&self) -> &dyn Any {
self
}

fn render_diff(&self, previous: Option<&dyn Any>) -> Option<Box<dyn ContextualUserFragment>> {
let previous = match previous {
Some(previous) => {
let Some(previous) = previous.downcast_ref::<S>() else {
unreachable!("world-state section type must match its type ID");
};
Some(previous)
fn snapshot(&self) -> Option<Value> {
let mut snapshot = match serde_json::to_value(WorldStateSection::snapshot(self)) {
Ok(snapshot) => snapshot,
Err(err) => {
tracing::error!(
section_id = S::ID,
%err,
"failed to serialize world-state section snapshot"
);
return None;
}
None => None,
};
WorldStateSection::render_diff(self, previous)
remove_null_object_fields(&mut snapshot);
Some(snapshot)
}

fn render_diff(&self, previous: Option<&Value>) -> Option<Box<dyn ContextualUserFragment>> {
let previous = previous.and_then(|previous| {
serde_json::from_value::<S::Snapshot>(previous.clone())
.inspect_err(|err| {
tracing::warn!(
section_id = S::ID,
%err,
"failed to restore world-state section snapshot"
);
})
.ok()
});
WorldStateSection::render_diff(self, previous.as_ref())
}
}

/// A typed portion of the state visible to the model.
///
/// Implementations own how their current state is rendered relative to an
/// earlier value of the same section type. A missing previous value requests
/// the section's complete current representation.
pub(crate) trait WorldStateSection: Any + Send + Sync {
fn render_diff(&self, previous: Option<&Self>) -> Option<Box<dyn ContextualUserFragment>>;
/// earlier snapshot of the same section. `ID` is persisted in rollouts and
/// must remain stable. `Snapshot` should contain only the comparison data
/// needed to decide what the model must be told next.
pub(crate) trait WorldStateSection: Send + Sync + 'static {
const ID: &'static str;
type Snapshot: DeserializeOwned + Serialize;

fn snapshot(&self) -> Self::Snapshot;

fn render_diff(
&self,
previous: Option<&Self::Snapshot>,
) -> Option<Box<dyn ContextualUserFragment>>;
}

/// A snapshot of the model-visible world with one section per concrete type.
/// Live model-visible state, keyed by the same stable section IDs used in rollouts.
#[derive(Default)]
pub(crate) struct WorldState {
sections: IndexMap<TypeId, Box<dyn ErasedWorldStateSection>>,
sections: IndexMap<&'static str, Box<dyn ErasedWorldStateSection>>,
}

/// Compact comparison state for each model-visible world-state section.
#[derive(Clone, Debug, Default, PartialEq, Serialize, serde::Deserialize)]
#[serde(transparent)]
pub(crate) struct WorldStateSnapshot {
sections: BTreeMap<String, Value>,
}

impl fmt::Debug for WorldState {
Expand All @@ -58,23 +90,55 @@ impl fmt::Debug for WorldState {

impl WorldState {
pub(crate) fn add_section<S: WorldStateSection>(&mut self, section: S) {
self.sections.insert(TypeId::of::<S>(), Box::new(section));
let id = S::ID;
assert!(
Comment thread
sayan-oai marked this conversation as resolved.
!self.sections.contains_key(id),
"duplicate world-state section ID: {id}"
);
self.sections.insert(id, Box::new(section));
}

pub(crate) fn snapshot(&self) -> WorldStateSnapshot {
WorldStateSnapshot {
sections: self
.sections
.iter()
.filter_map(|(id, section)| {
section
.snapshot()
.map(|snapshot| ((*id).to_string(), snapshot))
})
.collect(),
}
}

pub(crate) fn render_full(&self) -> Vec<Box<dyn ContextualUserFragment>> {
self.render_diff(&Self::default())
self.render_diff(&WorldStateSnapshot::default())
}

pub(crate) fn render_diff(&self, previous: &Self) -> Vec<Box<dyn ContextualUserFragment>> {
pub(crate) fn render_diff(
&self,
previous: &WorldStateSnapshot,
) -> Vec<Box<dyn ContextualUserFragment>> {
self.sections
.iter()
.filter_map(|(type_id, section)| {
let previous = previous
.sections
.get(type_id)
.map(|section| section.as_any());
section.render_diff(previous)
})
.filter_map(|(id, section)| section.render_diff(previous.sections.get(*id)))
.collect()
}
}

fn remove_null_object_fields(value: &mut Value) {
// RFC 7386 reserves object-valued nulls for deletion, but arrays are replaced whole.
match value {
Value::Object(values) => {
values.retain(|_, value| !value.is_null());
values.values_mut().for_each(remove_null_object_fields);
}
Value::Array(_) => {}
Value::Null | Value::Bool(_) | Value::Number(_) | Value::String(_) => {}
}
}

#[cfg(test)]
#[path = "world_state_tests.rs"]
mod tests;
Loading
Loading