Files
buzz/crates/buzz-cli/src/agent_management.rs

278 lines
8.8 KiB
Rust

//! Owner-reviewed agent draft requests published through Buzz observer frames.
use buzz_core::observer::{encrypt_observer_payload, OBSERVER_FRAME_TELEMETRY};
use nostr::{Event, Keys, PublicKey};
use serde::Serialize;
use crate::error::CliError;
const REQUEST_KIND: &str = "agent_management_request";
const MAX_NAME_CHARS: usize = 120;
const MAX_PROMPT_CHARS: usize = 20_000;
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct CreateAgentDraft {
pub channel_id: String,
pub display_name: String,
pub system_prompt: String,
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct UpdateAgentDraft {
pub channel_id: String,
pub agent_name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub display_name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub system_prompt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub runtime: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub provider: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub respond_to: Option<String>,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
struct ManagementRequest<T> {
#[serde(rename = "type")]
request_type: &'static str,
action: &'static str,
request_id: String,
request: T,
}
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
struct ObserverEvent<T> {
seq: u64,
timestamp: String,
kind: &'static str,
agent_index: Option<usize>,
channel_id: Option<String>,
session_id: Option<String>,
turn_id: Option<String>,
payload: ManagementRequest<T>,
}
#[derive(Debug)]
pub struct BuiltDraftRequest {
pub event: Event,
pub request_id: String,
pub action: &'static str,
}
fn required(value: String, label: &str, max: usize) -> Result<String, CliError> {
let value = value.trim();
if value.is_empty() {
return Err(CliError::Usage(format!("{label} is required")));
}
if value.chars().count() > max {
return Err(CliError::Usage(format!(
"{label} is too long (max {max} characters)"
)));
}
Ok(value.to_owned())
}
fn optional(value: Option<String>, label: &str) -> Result<Option<String>, CliError> {
value.map(|value| required(value, label, 300)).transpose()
}
fn build<T: Serialize>(
keys: &Keys,
owner: &PublicKey,
channel_id: String,
action: &'static str,
request: T,
) -> Result<BuiltDraftRequest, CliError> {
let request_id = uuid::Uuid::new_v4().to_string();
let payload = ObserverEvent {
seq: 0,
timestamp: chrono::Utc::now().to_rfc3339(),
kind: REQUEST_KIND,
agent_index: None,
channel_id: Some(channel_id),
session_id: None,
turn_id: None,
payload: ManagementRequest {
request_type: REQUEST_KIND,
action,
request_id: request_id.clone(),
request,
},
};
let encrypted = encrypt_observer_payload(keys, owner, &payload)
.map_err(|error| CliError::Other(format!("could not encrypt draft request: {error}")))?;
let event = buzz_sdk::build_agent_observer_frame(
&owner.to_hex(),
&keys.public_key().to_hex(),
OBSERVER_FRAME_TELEMETRY,
&encrypted,
)
.map_err(|error| CliError::Other(format!("could not build draft request: {error}")))?
.sign_with_keys(keys)
.map_err(|error| CliError::Other(format!("could not sign draft request: {error}")))?;
Ok(BuiltDraftRequest {
event,
request_id,
action,
})
}
pub fn build_create(
keys: &Keys,
owner: &PublicKey,
draft: CreateAgentDraft,
) -> Result<BuiltDraftRequest, CliError> {
let channel_id = required(draft.channel_id, "channel", 128)?;
uuid::Uuid::parse_str(&channel_id)
.map_err(|_| CliError::Usage(format!("invalid channel UUID: {channel_id}")))?;
let request = CreateAgentDraft {
channel_id: channel_id.clone(),
display_name: required(draft.display_name, "display name", MAX_NAME_CHARS)?,
system_prompt: required(draft.system_prompt, "system prompt", MAX_PROMPT_CHARS)?,
};
build(keys, owner, channel_id, "create", request)
}
pub fn build_update(
keys: &Keys,
owner: &PublicKey,
draft: UpdateAgentDraft,
) -> Result<BuiltDraftRequest, CliError> {
let channel_id = required(draft.channel_id, "channel", 128)?;
uuid::Uuid::parse_str(&channel_id)
.map_err(|_| CliError::Usage(format!("invalid channel UUID: {channel_id}")))?;
let respond_to = optional(draft.respond_to, "respond-to")?;
if respond_to
.as_deref()
.is_some_and(|value| value != "owner-only" && value != "anyone")
{
return Err(CliError::Usage(
"respond-to must be owner-only or anyone".into(),
));
}
let request = UpdateAgentDraft {
channel_id: channel_id.clone(),
agent_name: required(draft.agent_name, "agent name", MAX_NAME_CHARS)?,
display_name: optional(draft.display_name, "display name")?,
system_prompt: draft
.system_prompt
.map(|value| required(value, "system prompt", MAX_PROMPT_CHARS))
.transpose()?,
runtime: optional(draft.runtime, "runtime")?,
provider: optional(draft.provider, "provider")?,
model: optional(draft.model, "model")?,
respond_to,
};
if request.display_name.is_none()
&& request.system_prompt.is_none()
&& request.runtime.is_none()
&& request.provider.is_none()
&& request.model.is_none()
&& request.respond_to.is_none()
{
return Err(CliError::Usage(
"include at least one field to update".into(),
));
}
build(keys, owner, channel_id, "update", request)
}
#[cfg(test)]
mod tests {
use super::*;
use buzz_core::observer::{decrypt_observer_payload, OBSERVER_AGENT_TAG, OBSERVER_FRAME_TAG};
const CHANNEL: &str = "7c07e659-3610-42f4-9a5e-1e9973c09da9";
#[test]
fn create_is_owner_encrypted_and_matches_desktop_contract() {
let agent = Keys::generate();
let owner = Keys::generate();
let built = build_create(
&agent,
&owner.public_key(),
CreateAgentDraft {
channel_id: CHANNEL.into(),
display_name: "Research helper".into(),
system_prompt: "Find sources.".into(),
},
)
.unwrap();
assert_eq!(built.event.kind.as_u16(), 24_200);
let tags: Vec<Vec<String>> = built
.event
.tags
.iter()
.map(|tag| tag.as_slice().to_vec())
.collect();
assert!(tags
.iter()
.any(|tag| tag == &["p", &owner.public_key().to_hex()]));
assert!(tags
.iter()
.any(|tag| tag == &[OBSERVER_AGENT_TAG, &agent.public_key().to_hex()]));
assert!(tags
.iter()
.any(|tag| tag == &[OBSERVER_FRAME_TAG, OBSERVER_FRAME_TELEMETRY]));
assert!(!tags
.iter()
.any(|tag| tag.first().map(String::as_str) == Some("h")));
let payload: serde_json::Value = decrypt_observer_payload(&owner, &built.event).unwrap();
assert_eq!(payload["kind"], REQUEST_KIND);
assert_eq!(payload["channelId"], CHANNEL);
assert_eq!(payload["payload"]["type"], REQUEST_KIND);
assert_eq!(payload["payload"]["action"], "create");
assert_eq!(
payload["payload"]["request"]["displayName"],
"Research helper"
);
assert!(payload["payload"]["request"].get("runtime").is_none());
assert!(payload["payload"]["request"].get("respondTo").is_none());
}
#[test]
fn update_requires_a_change() {
let error = build_update(
&Keys::generate(),
&Keys::generate().public_key(),
UpdateAgentDraft {
channel_id: CHANNEL.into(),
agent_name: "Scout".into(),
display_name: None,
system_prompt: None,
runtime: None,
provider: None,
model: None,
respond_to: None,
},
)
.unwrap_err();
assert!(error.to_string().contains("at least one field"));
}
#[test]
fn create_rejects_invalid_channel() {
let error = build_create(
&Keys::generate(),
&Keys::generate().public_key(),
CreateAgentDraft {
channel_id: "general".into(),
display_name: "Scout".into(),
system_prompt: "Help".into(),
},
)
.unwrap_err();
assert!(error.to_string().contains("invalid channel UUID"));
}
}