diff --git a/src/commands/pair/agent.rs b/src/commands/pair/agent.rs index c132368..8b76673 100644 --- a/src/commands/pair/agent.rs +++ b/src/commands/pair/agent.rs @@ -1,3 +1,4 @@ +use std::ffi::OsString; use std::path::PathBuf; use crate::error::{self, Result}; @@ -49,9 +50,11 @@ impl Agent { } /// The prompt goes in as a single argument. It carries no token — only the - /// path to one — so it is safe in `ps` and in shell history. - pub fn command(&self, prompt: &str) -> tokio::process::Command { + /// path to one — so it is safe in `ps` and in shell history. Extra args go + /// in front of it, matching the agents' `[options] [prompt]` grammar. + pub fn command(&self, prompt: &str, extra_args: &[OsString]) -> tokio::process::Command { let mut command = tokio::process::Command::new(self.binary()); + command.args(extra_args); match self { Agent::Claude | Agent::Codex => command.arg(prompt), Agent::Opencode => command.args(["--prompt", prompt]), @@ -205,8 +208,9 @@ mod tests { } /// The prompt is one argv element, never split and never a shell string. - fn argv(agent: Agent) -> Vec { - let command = agent.command("pair with me"); + fn argv(agent: Agent, extra_args: &[&str]) -> Vec { + let extra_args: Vec = extra_args.iter().map(OsString::from).collect(); + let command = agent.command("pair with me", &extra_args); std::iter::once(command.as_std().get_program()) .chain(command.as_std().get_args()) .map(|arg| arg.to_string_lossy().into_owned()) @@ -215,18 +219,34 @@ mod tests { #[test] fn claude_and_codex_take_the_prompt_as_their_first_argument() { - assert_eq!(argv(Agent::Claude), ["claude", "pair with me"]); - assert_eq!(argv(Agent::Codex), ["codex", "pair with me"]); + assert_eq!(argv(Agent::Claude, &[]), ["claude", "pair with me"]); + assert_eq!(argv(Agent::Codex, &[]), ["codex", "pair with me"]); } #[test] fn opencode_takes_the_prompt_behind_its_prompt_flag() { assert_eq!( - argv(Agent::Opencode), + argv(Agent::Opencode, &[]), ["opencode", "--prompt", "pair with me"] ); } + #[test] + fn extra_args_go_before_the_prompt() { + assert_eq!( + argv(Agent::Claude, &["--model", "opus"]), + ["claude", "--model", "opus", "pair with me"] + ); + assert_eq!( + argv(Agent::Codex, &["--model", "opus"]), + ["codex", "--model", "opus", "pair with me"] + ); + assert_eq!( + argv(Agent::Opencode, &["--model", "opus"]), + ["opencode", "--model", "opus", "--prompt", "pair with me"] + ); + } + #[test] fn an_explicit_agent_is_used_when_it_is_ready() { let selected = select(Some(Agent::Opencode), given([READY, READY, READY])).unwrap(); diff --git a/src/commands/pair/mod.rs b/src/commands/pair/mod.rs index b468056..4432b17 100644 --- a/src/commands/pair/mod.rs +++ b/src/commands/pair/mod.rs @@ -3,6 +3,8 @@ mod prompt; mod session; mod target; +use std::ffi::OsString; + use clap::Args; use serde::Serialize; @@ -12,14 +14,14 @@ use crate::error::Result; use agent::{select, Agent}; use prompt::{build_prompt, write_token}; use session::Sessions; -use target::{resolve_dataset_editor, resolve_editor, PairTarget}; +use target::{resolve, PairTarget}; #[derive(Args, Debug, Serialize)] #[command(about = "Pair an agent CLI with a workspace's marimo notebook")] pub struct Pair { /// The workspace, or dataset with --dataset, to pair on, as - /// "{owner}/{slug}" with an optional "@{version}". Defaults to the newest - /// draft version. + /// "{owner}/{slug}" with an optional "@{version}", or a workspace version, + /// dataset or runner id. Defaults to the newest draft version. target: String, /// Pair on a dataset's notebook instead of a workspace's #[arg(long)] @@ -27,6 +29,9 @@ pub struct Pair { /// The notebook to open, defaulting to the workspace's overview notebook #[arg(long)] notebook: Option, + /// The marimo session to target, defaulting to the one live session + #[arg(long)] + session: Option, /// Pair with Claude Code instead of the first agent found #[arg(long, group = "agent")] claude: bool, @@ -42,6 +47,19 @@ pub struct Pair { /// Print the prompt instead of launching an agent #[arg(long, conflicts_with = "agent")] prompt_only: bool, + /// Extra arguments passed through to the agent command + #[arg(last = true, conflicts_with = "prompt_only")] + #[serde(serialize_with = "lossy_strings")] + agent_args: Vec, +} + +/// An `OsString` serializes as platform-tagged bytes, which is noise in the +/// command context sent to Sentry; the text is what anyone reading it wants. +fn lossy_strings( + args: &[OsString], + serializer: S, +) -> Result { + serializer.collect_seq(args.iter().map(|arg| arg.to_string_lossy())) } impl Pair { @@ -70,11 +88,7 @@ pub async fn pair(args: Pair, global: GlobalArgs) -> Result<()> { .spinner() .with_message(format!("Resolving the editor for {target}")); let client = global.graphql_client().await?; - let editor = if args.dataset { - resolve_dataset_editor(&client, &target, args.notebook).await? - } else { - resolve_editor(&client, &target, args.notebook).await? - }; + let editor = resolve(&client, &target, args.dataset, args.notebook).await?; pb.set_message(format!("Editor for {target} is {}", editor.phase)); let (token_dir, token_path) = write_token(&editor.token)?; @@ -85,7 +99,8 @@ pub async fn pair(args: Pair, global: GlobalArgs) -> Result<()> { // A session only exists while the notebook is open in a browser, so open it // — but not if the user already has it open. let sessions = Sessions::new(&editor, global.allow_insecure_host)?; - if !sessions.is_ready().await { + let mut live = sessions.live().await; + if live.is_empty() { if args.no_open { pb.println(format!("Please open {editor_page} to start the notebook")); } else { @@ -97,14 +112,19 @@ pub async fn pair(args: Pair, global: GlobalArgs) -> Result<()> { } } pb.set_message("Waiting for the notebook to connect"); - sessions.wait(&pb, &editor_page).await?; + live = sessions.wait(&pb, &editor_page).await?; } - let prompt = build_prompt(&editor, &token_path, &editor_page); + let session = session::choose(args.session.as_deref(), &live)?; + let prompt = build_prompt(&editor, &token_path, &editor_page, session); match agent { Some(agent) => { pb.finish_with_message(format!("Launching {}", agent.display_name())); - agent.command(&prompt).spawn()?.wait().await?; + agent + .command(&prompt, &args.agent_args) + .spawn()? + .wait() + .await?; // The agent is done with the token now. drop(token_dir); } @@ -143,6 +163,27 @@ mod tests { assert_eq!(args.agent(), None); assert!(!args.no_open); assert!(!args.prompt_only); + assert!(args.agent_args.is_empty()); + } + + #[test] + fn parses_extra_agent_args_after_a_double_dash() { + let args = parse(&["alice/ws", "--", "--model", "opus"]).unwrap(); + assert_eq!(args.agent_args, ["--model", "opus"]); + } + + /// The command is sent to Sentry as JSON; an `OsString` would arrive as + /// platform-tagged bytes rather than the text it holds. + #[test] + fn agent_args_serialize_as_strings() { + let args = parse(&["alice/ws", "--", "--model", "opus"]).unwrap(); + let json = serde_json::to_value(&args).unwrap(); + assert_eq!(json["agent_args"], serde_json::json!(["--model", "opus"])); + } + + #[test] + fn rejects_extra_agent_args_with_prompt_only() { + assert!(parse(&["alice/ws", "--prompt-only", "--", "--model", "opus"]).is_err()); } #[test] @@ -158,6 +199,12 @@ mod tests { assert!(parse(&["alice/ds", "--dataset"]).unwrap().dataset); } + #[test] + fn parses_the_session_flag() { + let args = parse(&["alice/ws", "--session", "s_1"]).unwrap(); + assert_eq!(args.session.as_deref(), Some("s_1")); + } + #[test] fn rejects_two_agent_flags() { assert!(parse(&["alice/ws", "--claude", "--codex"]).is_err()); diff --git a/src/commands/pair/prompt.rs b/src/commands/pair/prompt.rs index 3229b7d..96e7dba 100644 --- a/src/commands/pair/prompt.rs +++ b/src/commands/pair/prompt.rs @@ -14,25 +14,44 @@ fn quote(value: impl std::fmt::Display) -> String { format!("'{}'", value.to_string().replace('\'', r"'\''")) } -pub fn build_prompt(editor: &PairEditor, token_path: &Path, editor_page: &Url) -> String { +/// `session` is the one session to target, when there is exactly one; the +/// scripts resolve it themselves otherwise. +pub fn build_prompt( + editor: &PairEditor, + token_path: &Path, + editor_page: &Url, + session: Option<&str>, +) -> String { + let session_flag = session + .map(|id| format!(" --session {}", quote(id))) + .unwrap_or_default(); + let execute_cmd = format!( + "execute-code.sh --url {}{session_flag}", + quote(&editor.base_url) + ); + // The id is current now, but marimo renames a session when the browser + // reconnects, so the agent needs a way out. + let session_hint = if session.is_some() { + " If the script reports the session id is stale, drop --session and try again." + } else { + "" + }; format!( "Use the /marimo-pair skill to pair-program on a running marimo notebook. Connect to the notebook at: {base_url} -Use `execute-code.sh --url {quoted_url}` from the marimo-pair skill to execute code in the \ -notebook. +Use `{execute_cmd}` from the marimo-pair skill to execute code in the notebook. -An auth token is stored at {token_path}. Pass it via `execute-code.sh --url {quoted_url} \ +An auth token is stored at {token_path}. Pass it via `{execute_cmd} \ --token \"$(cat {quoted_token_path})\"`. The notebook must be open in a browser for a session to exist. If the server reports no \ -active sessions, ask the user to open {editor_page} and then try again. +active sessions, ask the user to open {editor_page} and then try again.{session_hint} Once you are connected, send a fun toast (mo.status.toast(...)) to the user inside marimo \ letting them know you're ready to pair.", base_url = editor.base_url, - quoted_url = quote(&editor.base_url), token_path = token_path.display(), quoted_token_path = quote(token_path.display()), editor_page = editor_page, @@ -96,17 +115,27 @@ fn open_private(path: &Path) -> Result { mod tests { use super::*; - #[test] - fn commands_carry_a_quoted_url_and_token_path() { - let editor = PairEditor { - base_url: Url::parse("http://host/runner/it's/").unwrap(), + fn editor(base_url: &str) -> PairEditor { + PairEditor { + base_url: Url::parse(base_url).unwrap(), token: "unused".into(), phase: "READY".into(), editor_page_id: "workspace-id".into(), - }; - let editor_page = Url::parse("https://aqora.io/workspaces/workspace-id/edit").unwrap(); + } + } - let prompt = build_prompt(&editor, Path::new("/tmp/it's dir/token.txt"), &editor_page); + fn editor_page() -> Url { + Url::parse("https://aqora.io/workspaces/workspace-id/edit").unwrap() + } + + #[test] + fn commands_carry_a_quoted_url_and_token_path() { + let prompt = build_prompt( + &editor("http://host/runner/it's/"), + Path::new("/tmp/it's dir/token.txt"), + &editor_page(), + None, + ); assert!( prompt.contains(r#"--url 'http://host/runner/it'\''s/'"#), @@ -116,6 +145,21 @@ mod tests { prompt.contains(r#"cat '/tmp/it'\''s dir/token.txt'"#), "{prompt}" ); + assert!(!prompt.contains("--session"), "{prompt}"); + } + + #[test] + fn commands_carry_the_session_when_it_is_known() { + let prompt = build_prompt( + &editor("http://host/runner/abc/"), + Path::new("/tmp/token.txt"), + &editor_page(), + Some("s_1"), + ); + + // Both the bare command and the one with the token target the session. + assert_eq!(prompt.matches("--session 's_1'").count(), 2, "{prompt}"); + assert!(prompt.contains("stale"), "{prompt}"); } #[test] diff --git a/src/commands/pair/session.rs b/src/commands/pair/session.rs index fe80dd0..653cf89 100644 --- a/src/commands/pair/session.rs +++ b/src/commands/pair/session.rs @@ -41,20 +41,20 @@ impl Sessions { }) } - /// Whether a session is live right now. A runner that is still starting + /// The ids of the sessions live right now. A runner that is still starting /// refuses connections and answers errors, so anything that is not a clear - /// "yes" counts as "not yet" — the caller decides how long to keep asking. - pub async fn is_ready(&self) -> bool { + /// answer counts as "none yet" — the caller decides how long to keep asking. + pub async fn live(&self) -> Vec { match self.query().await { - Ok(ready) => ready, + Ok(live) => live, Err(err) => { tracing::debug!("Could not read sessions from {}: {err}", self.url); - false + Vec::new() } } } - async fn query(&self) -> Result { + async fn query(&self) -> Result> { let response = self .client .get(self.url.clone()) @@ -62,15 +62,16 @@ impl Sessions { .send() .await? .error_for_status()?; - has_session(&response.bytes().await?) + session_ids(&response.bytes().await?) } /// Poll until the notebook connects, or give up. - pub async fn wait(&self, pb: &ProgressBar, editor_page: &Url) -> Result<()> { + pub async fn wait(&self, pb: &ProgressBar, editor_page: &Url) -> Result> { let deadline = tokio::time::Instant::now() + TIMEOUT; loop { - if self.is_ready().await { - return Ok(()); + let live = self.live().await; + if !live.is_empty() { + return Ok(live); } // Give up rather than sleeping through the deadline first. if tokio::time::Instant::now() + POLL_INTERVAL >= deadline { @@ -91,9 +92,28 @@ impl Sessions { /// `/api/sessions` answers an object keyed by session id, so an empty object /// means no notebook is open. -fn has_session(body: &[u8]) -> Result { +fn session_ids(body: &[u8]) -> Result> { let sessions: serde_json::Map = serde_json::from_slice(body)?; - Ok(!sessions.is_empty()) + Ok(sessions.into_iter().map(|(id, _)| id).collect()) +} + +/// The session the prompt targets: the one asked for, so long as it is live, +/// else the one live session. Several live and none asked for means the +/// notebook is open in more than one tab with no telling which the user means, +/// so the agent picks. +pub fn choose<'a>(asked: Option<&'a str>, live: &'a [String]) -> Result> { + match (asked, live) { + (Some(id), _) if live.iter().any(|s| s == id) => Ok(Some(id)), + (Some(id), _) => Err(error::user( + &format!("Session {id} is not open on the runner"), + &format!( + "The live sessions are: {}. Pass one of those, or drop --session.", + live.join(", ") + ), + )), + (None, [only]) => Ok(Some(only)), + (None, _) => Ok(None), + } } #[cfg(test)] @@ -147,7 +167,7 @@ mod tests { let (url, served) = serve_once(r#"{"s_1": {"path": "overview.py"}}"#).await; let sessions = Sessions::new(&editor(url), false).unwrap(); - assert!(sessions.is_ready().await); + assert_eq!(sessions.live().await, ["s_1"]); let request = served.await.unwrap().to_lowercase(); assert!( @@ -178,34 +198,60 @@ mod tests { }); let sessions = Sessions::new(&editor(url), false).unwrap(); - let ready = tokio::time::timeout(Duration::from_secs(15), sessions.is_ready()) + let live = tokio::time::timeout(Duration::from_secs(15), sessions.live()) .await - .expect("is_ready never gave up"); + .expect("live never gave up"); - assert!(!ready); + assert!(live.is_empty()); } #[tokio::test] - async fn is_not_ready_when_the_runner_reports_no_sessions() { + async fn nothing_is_live_when_the_runner_reports_no_sessions() { let (url, _served) = serve_once("{}").await; let sessions = Sessions::new(&editor(url), false).unwrap(); - assert!(!sessions.is_ready().await); + assert!(sessions.live().await.is_empty()); } #[test] - fn no_session_when_the_map_is_empty() { - assert!(!has_session(b"{}").unwrap()); + fn no_session_ids_when_the_map_is_empty() { + assert!(session_ids(b"{}").unwrap().is_empty()); } #[test] - fn a_session_is_ready_when_the_map_has_an_entry() { + fn session_ids_are_the_maps_keys() { let body = br#"{"s_1234": {"path": "/notebooks/overview.py"}}"#; - assert!(has_session(body).unwrap()); + assert_eq!(session_ids(body).unwrap(), ["s_1234"]); } #[test] fn errors_on_a_body_that_is_not_json() { - assert!(has_session(b"not marimo").is_err()); + assert!(session_ids(b"not marimo").is_err()); + } + + #[test] + fn the_session_is_known_when_exactly_one_is_live() { + assert_eq!(choose(None, &["s_1".to_string()]).unwrap(), Some("s_1")); + } + + #[test] + fn no_session_is_known_with_none_or_several_live() { + assert_eq!(choose(None, &[]).unwrap(), None); + let two = ["s_1".to_string(), "s_2".to_string()]; + assert_eq!(choose(None, &two).unwrap(), None); + } + + #[test] + fn an_asked_for_session_is_used_when_it_is_live() { + let two = ["s_1".to_string(), "s_2".to_string()]; + assert_eq!(choose(Some("s_2"), &two).unwrap(), Some("s_2")); + } + + #[test] + fn an_asked_for_session_that_is_not_live_is_an_error() { + let err = choose(Some("s_9"), &["s_1".to_string()]).unwrap_err(); + assert!(err.is_user()); + assert!(err.to_string().contains("s_9"), "{err}"); + assert!(err.to_string().contains("s_1"), "{err}"); } } diff --git a/src/commands/pair/target.rs b/src/commands/pair/target.rs index 1e48e2b..9e311f5 100644 --- a/src/commands/pair/target.rs +++ b/src/commands/pair/target.rs @@ -6,6 +6,7 @@ use url::Url; use crate::{ error::{self, Result}, graphql_client::{custom_scalars::*, GraphQLClient}, + id::{Id, NodeType}, }; #[derive(GraphQLQuery)] @@ -40,24 +41,96 @@ pub struct DatasetPairEditor; )] pub struct DatasetVersionPairEditor; +#[derive(GraphQLQuery)] +#[graphql( + query_path = "src/graphql/workspace_version_pair_editor_by_id.graphql", + schema_path = "schema.graphql", + response_derives = "Debug" +)] +pub struct WorkspaceVersionPairEditorById; + +#[derive(GraphQLQuery)] +#[graphql( + query_path = "src/graphql/dataset_pair_editor_by_id.graphql", + schema_path = "schema.graphql", + response_derives = "Debug" +)] +pub struct DatasetPairEditorById; + +#[derive(GraphQLQuery)] +#[graphql( + query_path = "src/graphql/workspace_runner_pair_editor_by_id.graphql", + schema_path = "schema.graphql", + response_derives = "Debug" +)] +pub struct WorkspaceRunnerPairEditorById; + +/// What to pair on: a workspace or dataset by slug, or a workspace version, +/// dataset or runner by node id. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum PairTarget { + Slug(SlugTarget), + /// A `WorkspaceVersion` node id, kept as written for messages and queries. + WorkspaceVersionId(String), + /// A `Dataset` node id, likewise. + DatasetId(String), + /// A `WorkspaceRunner` node id, likewise. + RunnerId(String), +} + /// A workspace to pair on, written `owner/slug` with an optional `@version`. #[derive(Debug, Clone, PartialEq, Eq)] -pub struct PairTarget { +pub struct SlugTarget { pub owner: String, pub slug: String, pub version: Option, } -const TARGET_ADVICE: &str = "Expected a workspace like: {owner}/{workspace}[@{version}]"; +const TARGET_ADVICE: &str = "Expected a workspace like {owner}/{workspace}[@{version}], or a \ + workspace version, dataset or runner id"; +const SLUG_ADVICE: &str = "Expected a workspace like: {owner}/{workspace}[@{version}]"; impl FromStr for PairTarget { type Err = crate::error::Error; + fn from_str(input: &str) -> Result { + // Node ids are base64, so a slash can only mean a slug. + if input.contains('/') { + return input.parse().map(PairTarget::Slug); + } + let id = + Id::parse_node_id(input).map_err(|_| error::user("Malformed target", TARGET_ADVICE))?; + match id.ty { + NodeType::WorkspaceVersion => Ok(PairTarget::WorkspaceVersionId(input.to_string())), + NodeType::Dataset => Ok(PairTarget::DatasetId(input.to_string())), + NodeType::WorkspaceRunner => Ok(PairTarget::RunnerId(input.to_string())), + other => Err(error::user( + &format!("{input} is a {other} id"), + TARGET_ADVICE, + )), + } + } +} + +impl std::fmt::Display for PairTarget { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + PairTarget::Slug(slug) => slug.fmt(f), + PairTarget::WorkspaceVersionId(id) + | PairTarget::DatasetId(id) + | PairTarget::RunnerId(id) => f.write_str(id), + } + } +} + +impl FromStr for SlugTarget { + type Err = crate::error::Error; + fn from_str(input: &str) -> Result { let input = input.strip_prefix('@').unwrap_or(input); let (owner, rest) = input .split_once('/') - .ok_or_else(|| error::user("Malformed workspace", TARGET_ADVICE))?; + .ok_or_else(|| error::user("Malformed workspace", SLUG_ADVICE))?; // Only an `@` *after* the slash introduces a version, so a leading // `@owner` stays part of the owner. @@ -75,10 +148,10 @@ impl FromStr for PairTarget { }; if owner.is_empty() || slug.is_empty() { - return Err(error::user("Malformed workspace", TARGET_ADVICE)); + return Err(error::user("Malformed workspace", SLUG_ADVICE)); } - Ok(PairTarget { + Ok(SlugTarget { owner: owner.to_string(), slug: slug.to_string(), version, @@ -86,7 +159,7 @@ impl FromStr for PairTarget { } } -impl std::fmt::Display for PairTarget { +impl std::fmt::Display for SlugTarget { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write!(f, "{}/{}", self.owner, self.slug)?; if let Some(version) = &self.version { @@ -124,7 +197,7 @@ fn split_url_and_token(mut url: Url) -> Result<(Url, String)> { Ok((url, token)) } -fn no_editor(target: &PairTarget, published: bool) -> crate::error::Error { +fn no_editor(target: impl std::fmt::Display, published: bool) -> crate::error::Error { if published { error::user( &format!("{target} is published and has no editor"), @@ -140,7 +213,7 @@ fn no_editor(target: &PairTarget, published: bool) -> crate::error::Error { } } -fn no_draft(target: &PairTarget) -> crate::error::Error { +fn no_draft(target: impl std::fmt::Display) -> crate::error::Error { error::user( &format!("{target} has no draft version"), "Pairing edits a workspace's draft version. Create a draft version on aqora.io, \ @@ -148,7 +221,7 @@ fn no_draft(target: &PairTarget) -> crate::error::Error { ) } -fn cannot_edit(target: &PairTarget) -> crate::error::Error { +fn cannot_edit(target: impl std::fmt::Display) -> crate::error::Error { error::user( &format!("You cannot edit {target}"), "Pairing needs edit access to the version. Check that you are logged in as a user \ @@ -156,9 +229,32 @@ fn cannot_edit(target: &PairTarget) -> crate::error::Error { ) } -pub async fn resolve_editor( +/// Resolves whatever was asked for to its editor. An id already says what it +/// is, so `--dataset` only picks between the two things a slug can name — and +/// a runner may be a dataset's editor, so there the flag is neither right nor +/// wrong. +pub async fn resolve( client: &GraphQLClient, target: &PairTarget, + dataset: bool, + notebook: Option, +) -> Result { + match target { + PairTarget::Slug(slug) if dataset => resolve_dataset_editor(client, slug, notebook).await, + PairTarget::Slug(slug) => resolve_editor(client, slug, notebook).await, + PairTarget::WorkspaceVersionId(_) if dataset => Err(error::user( + &format!("{target} is a workspace version, not a dataset"), + "Drop --dataset to pair on a workspace version by id", + )), + PairTarget::WorkspaceVersionId(id) => resolve_version_by_id(client, id, notebook).await, + PairTarget::DatasetId(id) => resolve_dataset_by_id(client, id, notebook).await, + PairTarget::RunnerId(id) => resolve_runner_by_id(client, id, notebook).await, + } +} + +async fn resolve_editor( + client: &GraphQLClient, + target: &SlugTarget, notebook: Option, ) -> Result { match &target.version { @@ -169,7 +265,7 @@ pub async fn resolve_editor( async fn resolve_pinned( client: &GraphQLClient, - target: &PairTarget, + target: &SlugTarget, version: &semver::Version, notebook: Option, ) -> Result { @@ -216,7 +312,7 @@ async fn resolve_pinned( async fn resolve_draft( client: &GraphQLClient, - target: &PairTarget, + target: &SlugTarget, notebook: Option, ) -> Result { let workspace = client @@ -253,14 +349,14 @@ async fn resolve_draft( }) } -fn workspace_not_found(target: &PairTarget) -> crate::error::Error { +fn workspace_not_found(target: &SlugTarget) -> crate::error::Error { error::user( &format!("Workspace {}/{} not found", target.owner, target.slug), "Please double check the workspace on aqora.io", ) } -fn dataset_not_found(target: &PairTarget) -> crate::error::Error { +fn dataset_not_found(target: &SlugTarget) -> crate::error::Error { error::user( &format!("Dataset {}/{} not found", target.owner, target.slug), "Please double check the dataset on aqora.io", @@ -268,9 +364,9 @@ fn dataset_not_found(target: &PairTarget) -> crate::error::Error { } /// A dataset is edited through the workspace its version owns. -pub async fn resolve_dataset_editor( +async fn resolve_dataset_editor( client: &GraphQLClient, - target: &PairTarget, + target: &SlugTarget, notebook: Option, ) -> Result { match &target.version { @@ -281,7 +377,7 @@ pub async fn resolve_dataset_editor( async fn resolve_dataset_pinned( client: &GraphQLClient, - target: &PairTarget, + target: &SlugTarget, version: &semver::Version, notebook: Option, ) -> Result { @@ -341,7 +437,7 @@ async fn resolve_dataset_pinned( async fn resolve_dataset_draft( client: &GraphQLClient, - target: &PairTarget, + target: &SlugTarget, notebook: Option, ) -> Result { let dataset = client @@ -377,18 +473,162 @@ async fn resolve_dataset_draft( }) } +/// The id encodes its kind, so the platform answering with another kind is a +/// bug, not a typo. +fn wrong_kind(id: &str) -> crate::error::Error { + error::system( + &format!("{id} resolved to an unexpected kind of node"), + "The platform returned an unexpected node. Please report this.", + ) +} + +async fn resolve_version_by_id( + client: &GraphQLClient, + id: &str, + notebook: Option, +) -> Result { + let response = client + .send::(workspace_version_pair_editor_by_id::Variables { + id: id.to_string(), + notebook, + }) + .await?; + let version = match response.node { + workspace_version_pair_editor_by_id::WorkspaceVersionPairEditorByIdNode::WorkspaceVersion( + version, + ) => version, + _ => return Err(wrong_kind(id)), + }; + + if version.published_at.is_some() { + return Err(error::user( + &format!("{id} is published and has no editor"), + "Published versions are read-only. Pair on the workspace's draft version \ + instead, or create a new draft.", + )); + } + if !version.viewer_can_edit { + return Err(cannot_edit(id)); + } + + let editor = version.editor.ok_or_else(|| no_editor(id, false))?; + let (base_url, token) = split_url_and_token(editor.url)?; + + Ok(PairEditor { + base_url, + token, + phase: format!("{:?}", editor.phase), + editor_page_id: version.id, + }) +} + +async fn resolve_dataset_by_id( + client: &GraphQLClient, + id: &str, + notebook: Option, +) -> Result { + let response = client + .send::(dataset_pair_editor_by_id::Variables { + id: id.to_string(), + notebook, + }) + .await?; + let dataset = match response.node { + dataset_pair_editor_by_id::DatasetPairEditorByIdNode::Dataset(dataset) => dataset, + _ => return Err(wrong_kind(id)), + }; + + let draft = dataset + .versions + .nodes + .into_iter() + .next() + .ok_or_else(|| no_draft(id))?; + + let workspace = draft.workspace.ok_or_else(|| no_editor(id, false))?; + if !workspace.viewer_can_edit { + return Err(cannot_edit(id)); + } + + let editor = workspace.editor.ok_or_else(|| no_editor(id, false))?; + let (base_url, token) = split_url_and_token(editor.url)?; + + Ok(PairEditor { + base_url, + token, + phase: format!("{:?}", editor.phase), + editor_page_id: workspace.id, + }) +} + +/// A runner is the editor itself. Reading an editor runner already needs edit +/// access, so the platform answers a permission error rather than a runner the +/// viewer cannot use. +async fn resolve_runner_by_id( + client: &GraphQLClient, + id: &str, + notebook: Option, +) -> Result { + let response = client + .send::(workspace_runner_pair_editor_by_id::Variables { + id: id.to_string(), + notebook, + }) + .await?; + let runner = match response.node { + workspace_runner_pair_editor_by_id::WorkspaceRunnerPairEditorByIdNode::WorkspaceRunner( + runner, + ) => runner, + _ => return Err(wrong_kind(id)), + }; + + if !matches!( + runner.command, + workspace_runner_pair_editor_by_id::RunnerCommand::EDIT + ) { + return Err(error::user( + &format!("{id} is a {:?} runner, not an editor", runner.command), + "Pairing needs the workspace's editor runner. Pair on the workspace or version \ + instead.", + )); + } + + let (base_url, token) = split_url_and_token(runner.url)?; + + Ok(PairEditor { + base_url, + token, + phase: format!("{:?}", runner.phase), + // The edit page takes a version or a workspace; dataset runners only + // have the latter. + editor_page_id: runner.workspace_version_id.unwrap_or(runner.workspace.id), + }) +} + #[cfg(test)] mod tests { use super::*; use aqora_client::ClientOptions; + use uuid::Uuid; - use crate::graphql_client::unauthenticated_client; + use crate::{ + graphql_client::unauthenticated_client, + id::{Id, NodeType}, + }; fn parse(input: &str) -> PairTarget { input.parse().unwrap() } + fn node_id(ty: NodeType) -> String { + Id { + id: Uuid::from_u128(0x1234), + ty, + } + .to_node_id() + } + /// Whether a request has been read in full, so the canned answer is not /// written before the query has arrived. fn is_complete(request: &[u8]) -> bool { @@ -454,7 +694,7 @@ mod tests { async fn a_draft_target_uses_the_draft_versions_editor() { let (client, _server) = serve_graphql(WORKSPACE_WITH_DRAFT).await; - let editor = resolve_editor(&client, &parse("alice/ws"), None) + let editor = resolve(&client, &parse("alice/ws"), false, None) .await .unwrap(); @@ -473,7 +713,7 @@ mod tests { ) .await; - let err = resolve_editor(&client, &parse("alice/ws"), None) + let err = resolve(&client, &parse("alice/ws"), false, None) .await .unwrap_err(); @@ -493,7 +733,7 @@ mod tests { ) .await; - let err = resolve_editor(&client, &parse("alice/ws"), None) + let err = resolve(&client, &parse("alice/ws"), false, None) .await .unwrap_err(); @@ -515,7 +755,7 @@ mod tests { async fn a_dataset_target_uses_the_draft_versions_workspace_editor() { let (client, _server) = serve_graphql(DATASET_WITH_DRAFT).await; - let editor = resolve_dataset_editor(&client, &parse("alice/ds"), None) + let editor = resolve(&client, &parse("alice/ds"), true, None) .await .unwrap(); @@ -535,7 +775,7 @@ mod tests { ) .await; - let err = resolve_dataset_editor(&client, &parse("alice/ds"), None) + let err = resolve(&client, &parse("alice/ds"), true, None) .await .unwrap_err(); @@ -547,7 +787,7 @@ mod tests { async fn a_dataset_target_errors_when_the_dataset_does_not_exist() { let (client, _server) = serve_graphql(r#"{"data":{"datasetBySlug":null}}"#).await; - let err = resolve_dataset_editor(&client, &parse("alice/ds"), None) + let err = resolve(&client, &parse("alice/ds"), true, None) .await .unwrap_err(); @@ -570,7 +810,7 @@ mod tests { ) .await; - let err = resolve_dataset_editor(&client, &parse("alice/ds"), None) + let err = resolve(&client, &parse("alice/ds"), true, None) .await .unwrap_err(); @@ -582,7 +822,7 @@ mod tests { async fn a_pinned_dataset_target_rejects_a_prerelease_version() { let (client, _server) = serve_graphql(DATASET_WITH_DRAFT).await; - let err = resolve_dataset_editor(&client, &parse("alice/ds@1.2.3-beta.1"), None) + let err = resolve(&client, &parse("alice/ds@1.2.3-beta.1"), true, None) .await .unwrap_err(); @@ -602,7 +842,7 @@ mod tests { ) .await; - let editor = resolve_dataset_editor(&client, &parse("alice/ds@1.2.3"), None) + let editor = resolve(&client, &parse("alice/ds@1.2.3"), true, None) .await .unwrap(); @@ -624,7 +864,7 @@ mod tests { ) .await; - let err = resolve_editor(&client, &parse("alice/ws@1.2.3"), None) + let err = resolve(&client, &parse("alice/ws@1.2.3"), false, None) .await .unwrap_err(); @@ -632,12 +872,274 @@ mod tests { assert!(err.to_string().contains("edit"), "{err}"); } + const WORKSPACE_VERSION_NODE: &str = r#"{"data":{"node":{"__typename":"WorkspaceVersion", + "id":"version-id","publishedAt":null,"viewerCanEdit":true, + "editor":{"id":"r-version","phase":"READY", + "url":"http://localhost:8080/runner/version/?access_token=version-token"}}}}"#; + + #[tokio::test] + async fn a_workspace_version_id_uses_that_versions_editor() { + let (client, served) = serve_graphql(WORKSPACE_VERSION_NODE).await; + let id = node_id(NodeType::WorkspaceVersion); + + let editor = resolve(&client, &parse(&id), false, None).await.unwrap(); + + assert_eq!( + editor.base_url.as_str(), + "http://localhost:8080/runner/version/" + ); + assert_eq!(editor.token, "version-token"); + assert_eq!(editor.editor_page_id, "version-id"); + let request = served.await.unwrap(); + assert!(request.contains(&id), "{request}"); + } + + #[tokio::test] + async fn a_workspace_version_id_errors_when_the_version_is_published() { + let (client, _server) = serve_graphql( + r#"{"data":{"node":{"__typename":"WorkspaceVersion", + "id":"version-id","publishedAt":"2026-01-01T00:00:00Z","viewerCanEdit":true, + "editor":{"id":"r-version","phase":"READY", + "url":"http://localhost:8080/runner/version/?access_token=t"}}}}"#, + ) + .await; + + let err = resolve( + &client, + &parse(&node_id(NodeType::WorkspaceVersion)), + false, + None, + ) + .await + .unwrap_err(); + + assert!(err.is_user()); + assert!(err.to_string().contains("published"), "{err}"); + } + + #[tokio::test] + async fn a_workspace_version_id_errors_when_the_viewer_cannot_edit_the_version() { + let (client, _server) = serve_graphql( + r#"{"data":{"node":{"__typename":"WorkspaceVersion", + "id":"version-id","publishedAt":null,"viewerCanEdit":false, + "editor":{"id":"r-version","phase":"READY", + "url":"http://localhost:8080/runner/version/?access_token=t"}}}}"#, + ) + .await; + + let err = resolve( + &client, + &parse(&node_id(NodeType::WorkspaceVersion)), + false, + None, + ) + .await + .unwrap_err(); + + assert!(err.is_user()); + assert!(err.to_string().contains("edit"), "{err}"); + } + + #[tokio::test] + async fn a_workspace_version_id_errors_when_the_version_has_no_editor() { + let (client, _server) = serve_graphql( + r#"{"data":{"node":{"__typename":"WorkspaceVersion", + "id":"version-id","publishedAt":null,"viewerCanEdit":true,"editor":null}}}"#, + ) + .await; + + let err = resolve( + &client, + &parse(&node_id(NodeType::WorkspaceVersion)), + false, + None, + ) + .await + .unwrap_err(); + + assert!(err.is_user()); + assert!(err.to_string().contains("No editor"), "{err}"); + } + + #[tokio::test] + async fn a_workspace_version_id_rejects_the_dataset_flag() { + let (client, _server) = serve_graphql(WORKSPACE_VERSION_NODE).await; + + let err = resolve( + &client, + &parse(&node_id(NodeType::WorkspaceVersion)), + true, + None, + ) + .await + .unwrap_err(); + + assert!(err.is_user()); + assert!(err.to_string().contains("--dataset"), "{err}"); + } + + const DATASET_NODE: &str = r#"{"data":{"node":{"__typename":"Dataset", + "versions":{"nodes":[{"id":"dataset-version-id","workspace":{ + "id":"workspace-id","viewerCanEdit":true, + "editor":{"id":"r-dataset","phase":"READY", + "url":"http://localhost:8080/runner/dataset/?access_token=dataset-token"}}}]}}}}"#; + + #[tokio::test] + async fn a_dataset_id_uses_the_draft_versions_workspace_editor() { + let (client, served) = serve_graphql(DATASET_NODE).await; + let id = node_id(NodeType::Dataset); + + let editor = resolve(&client, &parse(&id), false, None).await.unwrap(); + + assert_eq!( + editor.base_url.as_str(), + "http://localhost:8080/runner/dataset/" + ); + assert_eq!(editor.token, "dataset-token"); + assert_eq!(editor.editor_page_id, "workspace-id"); + let request = served.await.unwrap(); + assert!(request.contains(&id), "{request}"); + } + + #[tokio::test] + async fn a_dataset_id_accepts_the_dataset_flag() { + let (client, _server) = serve_graphql(DATASET_NODE).await; + + let editor = resolve(&client, &parse(&node_id(NodeType::Dataset)), true, None) + .await + .unwrap(); + + assert_eq!(editor.editor_page_id, "workspace-id"); + } + + #[tokio::test] + async fn a_dataset_id_errors_when_the_dataset_has_no_draft_version() { + let (client, _server) = + serve_graphql(r#"{"data":{"node":{"__typename":"Dataset","versions":{"nodes":[]}}}}"#) + .await; + + let err = resolve(&client, &parse(&node_id(NodeType::Dataset)), false, None) + .await + .unwrap_err(); + + assert!(err.is_user()); + assert!(err.to_string().contains("has no draft version"), "{err}"); + } + + #[tokio::test] + async fn a_dataset_id_errors_when_the_viewer_cannot_edit_the_workspace() { + let (client, _server) = serve_graphql( + r#"{"data":{"node":{"__typename":"Dataset","versions":{"nodes":[ + {"id":"dataset-version-id","workspace":{ + "id":"workspace-id","viewerCanEdit":false, + "editor":{"id":"r-dataset","phase":"READY", + "url":"http://localhost:8080/runner/dataset/?access_token=t"}}} + ]}}}}"#, + ) + .await; + + let err = resolve(&client, &parse(&node_id(NodeType::Dataset)), false, None) + .await + .unwrap_err(); + + assert!(err.is_user()); + assert!(err.to_string().contains("edit"), "{err}"); + } + + const RUNNER_NODE: &str = r#"{"data":{"node":{"__typename":"WorkspaceRunner", + "id":"r-runner","command":"EDIT","phase":"READY", + "url":"http://localhost:8080/runner/abc/?access_token=runner-token", + "workspaceVersionId":"version-id","workspace":{"id":"workspace-id"}}}}"#; + + #[tokio::test] + async fn a_runner_id_uses_that_runner() { + let (client, served) = serve_graphql(RUNNER_NODE).await; + let id = node_id(NodeType::WorkspaceRunner); + + let editor = resolve(&client, &parse(&id), false, None).await.unwrap(); + + assert_eq!( + editor.base_url.as_str(), + "http://localhost:8080/runner/abc/" + ); + assert_eq!(editor.token, "runner-token"); + assert_eq!(editor.editor_page_id, "version-id"); + let request = served.await.unwrap(); + assert!(request.contains(&id), "{request}"); + } + + #[tokio::test] + async fn a_runner_id_without_a_version_edits_through_its_workspace() { + let (client, _server) = serve_graphql( + r#"{"data":{"node":{"__typename":"WorkspaceRunner", + "id":"r-runner","command":"EDIT","phase":"READY", + "url":"http://localhost:8080/runner/abc/?access_token=t", + "workspaceVersionId":null,"workspace":{"id":"workspace-id"}}}}"#, + ) + .await; + + let editor = resolve( + &client, + &parse(&node_id(NodeType::WorkspaceRunner)), + false, + None, + ) + .await + .unwrap(); + + assert_eq!(editor.editor_page_id, "workspace-id"); + } + + #[tokio::test] + async fn a_runner_id_errors_when_the_runner_is_not_an_editor() { + let (client, _server) = serve_graphql( + r#"{"data":{"node":{"__typename":"WorkspaceRunner", + "id":"r-runner","command":"RENDER","phase":"READY", + "url":"http://localhost:8080/runner/abc/?access_token=t", + "workspaceVersionId":"version-id","workspace":{"id":"workspace-id"}}}}"#, + ) + .await; + + let err = resolve( + &client, + &parse(&node_id(NodeType::WorkspaceRunner)), + false, + None, + ) + .await + .unwrap_err(); + + assert!(err.is_user()); + assert!(err.to_string().contains("not an editor"), "{err}"); + } + + /// A runner may be a dataset's editor, so the flag is neither right nor wrong. + #[tokio::test] + async fn a_runner_id_accepts_the_dataset_flag() { + let (client, _server) = serve_graphql(RUNNER_NODE).await; + + let editor = resolve( + &client, + &parse(&node_id(NodeType::WorkspaceRunner)), + true, + None, + ) + .await + .unwrap(); + + assert_eq!(editor.editor_page_id, "version-id"); + } + #[test] fn parses_owner_and_slug() { - let target = parse("alice/my-workspace"); - assert_eq!(target.owner, "alice"); - assert_eq!(target.slug, "my-workspace"); - assert_eq!(target.version, None); + assert_eq!( + parse("alice/my-workspace"), + PairTarget::Slug(SlugTarget { + owner: "alice".to_string(), + slug: "my-workspace".to_string(), + version: None, + }) + ); } #[test] @@ -647,15 +1149,52 @@ mod tests { #[test] fn parses_version() { - let target = parse("@alice/my-workspace@1.2.3"); - assert_eq!(target.owner, "alice"); - assert_eq!(target.slug, "my-workspace"); - assert_eq!(target.version, Some(semver::Version::new(1, 2, 3))); + assert_eq!( + parse("@alice/my-workspace@1.2.3"), + PairTarget::Slug(SlugTarget { + owner: "alice".to_string(), + slug: "my-workspace".to_string(), + version: Some(semver::Version::new(1, 2, 3)), + }) + ); + } + + #[test] + fn parses_a_workspace_version_id() { + let id = node_id(NodeType::WorkspaceVersion); + assert_eq!(parse(&id), PairTarget::WorkspaceVersionId(id.clone())); + } + + #[test] + fn parses_a_dataset_id() { + let id = node_id(NodeType::Dataset); + assert_eq!(parse(&id), PairTarget::DatasetId(id.clone())); + } + + #[test] + fn parses_a_workspace_runner_id() { + let id = node_id(NodeType::WorkspaceRunner); + assert_eq!(parse(&id), PairTarget::RunnerId(id.clone())); + } + + #[test] + fn rejects_an_id_of_another_kind() { + let err = node_id(NodeType::ProviderJob) + .parse::() + .unwrap_err(); + assert!(err.is_user()); + assert!(err.to_string().contains("ProviderJob"), "{err}"); } #[test] fn round_trips_through_display() { - for input in ["alice/my-workspace", "alice/my-workspace@1.2.3"] { + for input in [ + "alice/my-workspace", + "alice/my-workspace@1.2.3", + node_id(NodeType::WorkspaceVersion).as_str(), + node_id(NodeType::Dataset).as_str(), + node_id(NodeType::WorkspaceRunner).as_str(), + ] { assert_eq!(parse(input).to_string(), input); } } diff --git a/src/graphql/dataset_pair_editor_by_id.graphql b/src/graphql/dataset_pair_editor_by_id.graphql new file mode 100644 index 0000000..8dec93e --- /dev/null +++ b/src/graphql/dataset_pair_editor_by_id.graphql @@ -0,0 +1,21 @@ +query DatasetPairEditorById($id: ID!, $notebook: String) { + node(id: $id) { + __typename + ... on Dataset { + versions(first: 1, filters: { published: false }) { + nodes { + id + workspace { + id + viewerCanEdit: viewerCan(action: UPDATE_WORKSPACE) + editor { + id + phase + url(notebook: $notebook) + } + } + } + } + } + } +} diff --git a/src/graphql/workspace_runner_pair_editor_by_id.graphql b/src/graphql/workspace_runner_pair_editor_by_id.graphql new file mode 100644 index 0000000..31383c8 --- /dev/null +++ b/src/graphql/workspace_runner_pair_editor_by_id.graphql @@ -0,0 +1,15 @@ +query WorkspaceRunnerPairEditorById($id: ID!, $notebook: String) { + node(id: $id) { + __typename + ... on WorkspaceRunner { + id + command + phase + url(notebook: $notebook) + workspaceVersionId + workspace { + id + } + } + } +} diff --git a/src/graphql/workspace_version_pair_editor_by_id.graphql b/src/graphql/workspace_version_pair_editor_by_id.graphql new file mode 100644 index 0000000..0a6f15c --- /dev/null +++ b/src/graphql/workspace_version_pair_editor_by_id.graphql @@ -0,0 +1,15 @@ +query WorkspaceVersionPairEditorById($id: ID!, $notebook: String) { + node(id: $id) { + __typename + ... on WorkspaceVersion { + id + publishedAt + viewerCanEdit: viewerCan(action: UPDATE_WORKSPACE_VERSION) + editor { + id + phase + url(notebook: $notebook) + } + } + } +} diff --git a/src/id.rs b/src/id.rs index 42a3859..0189aad 100644 --- a/src/id.rs +++ b/src/id.rs @@ -10,6 +10,9 @@ pub enum NodeType { ProjectVersionFile, ProviderModel, ProviderJob, + WorkspaceVersion, + Dataset, + WorkspaceRunner, } impl FromStr for NodeType { @@ -23,6 +26,9 @@ impl FromStr for NodeType { "ProjectVersionFile" => Ok(NodeType::ProjectVersionFile), "ProviderModel" => Ok(NodeType::ProviderModel), "ProviderJob" => Ok(NodeType::ProviderJob), + "WorkspaceVersion" => Ok(NodeType::WorkspaceVersion), + "Dataset" => Ok(NodeType::Dataset), + "WorkspaceRunner" => Ok(NodeType::WorkspaceRunner), _ => Err(format!("Unknown node kind: {}", s)), } } @@ -37,6 +43,9 @@ impl fmt::Display for NodeType { NodeType::ProjectVersionFile => write!(f, "ProjectVersionFile"), NodeType::ProviderModel => write!(f, "ProviderModel"), NodeType::ProviderJob => write!(f, "ProviderJob"), + NodeType::WorkspaceVersion => write!(f, "WorkspaceVersion"), + NodeType::Dataset => write!(f, "Dataset"), + NodeType::WorkspaceRunner => write!(f, "WorkspaceRunner"), } } }