diff --git a/TASKS.md b/TASKS.md index f8e1367..9279353 100644 --- a/TASKS.md +++ b/TASKS.md @@ -2,19 +2,19 @@ ## Current -### `feat/workspace-foundation` +### `feat/agent-publishing` Status: completed DoD: - -- добавлены `workspaces` и `workspace_id` в storage model -- `admin-api` переведен на `workspace-scoped` routes для operations и auth profiles -- тесты и registry migration проходят для workspace foundation +- можно создать `agent` в рамках `workspace` +- можно привязать published operations к `agent` +- `mcp-server` отдает tools в контексте `workspace + agent` +- один `agent` видит только свой curated toolset ## Next -- `feat/agent-publishing` +- `feat/platform-access` ## Backlog diff --git a/apps/admin-api/src/app.rs b/apps/admin-api/src/app.rs index c5a08f7..b2dc766 100644 --- a/apps/admin-api/src/app.rs +++ b/apps/admin-api/src/app.rs @@ -5,6 +5,10 @@ use axum::{ use crate::{ routes::{ + agents::{ + create_agent, get_agent, get_agent_version, list_agents, publish_agent, + save_agent_bindings, + }, auth_profiles::{create_auth_profile, get_auth_profile, list_auth_profiles}, operations::{ create_operation, create_version, export_operation, generate_draft, get_operation, @@ -57,6 +61,14 @@ pub fn build_app(state: AppState) -> Router { post(generate_draft), ) .route("/operations/{operation_id}/export", get(export_operation)) + .route("/agents", get(list_agents).post(create_agent)) + .route("/agents/{agent_id}", get(get_agent)) + .route( + "/agents/{agent_id}/versions/{version}", + get(get_agent_version), + ) + .route("/agents/{agent_id}/bindings", post(save_agent_bindings)) + .route("/agents/{agent_id}/publish", post(publish_agent)) .route( "/auth-profiles", get(list_auth_profiles).post(create_auth_profile), @@ -170,6 +182,86 @@ mod tests { assert_eq!(test_run["response_preview"]["id"], "lead_123"); } + #[tokio::test] + async fn creates_binds_and_publishes_agent() { + let registry = test_registry().await; + let storage_root = test_storage_root("agent_lifecycle"); + let upstream_base_url = spawn_upstream_server().await; + let base_url = spawn_admin_api(build_test_app(registry, storage_root)).await; + let client = reqwest::Client::new(); + + let operation = client + .post(format!("{base_url}/operations")) + .json(&test_operation_payload( + &upstream_base_url, + "crm_create_lead_agent", + )) + .send() + .await + .unwrap() + .json::() + .await + .unwrap(); + let operation_id = operation["operation_id"].as_str().unwrap().to_owned(); + + client + .post(format!("{base_url}/operations/{operation_id}/publish")) + .json(&json!({ "version": 1 })) + .send() + .await + .unwrap(); + + let agent = client + .post(format!("{base_url}/agents")) + .json(&json!({ + "slug": "sales-assistant", + "display_name": "Sales Assistant", + "description": "Curated sales toolset", + "instructions": {}, + "tool_selection_policy": {} + })) + .send() + .await + .unwrap() + .json::() + .await + .unwrap(); + let agent_id = agent["agent_id"].as_str().unwrap().to_owned(); + + let bindings = client + .post(format!("{base_url}/agents/{agent_id}/bindings")) + .json(&json!([ + { + "operation_id": operation_id, + "operation_version": 1, + "tool_name": "crm_create_lead_agent", + "tool_title": "Create Lead", + "enabled": true + } + ])) + .send() + .await + .unwrap() + .json::() + .await + .unwrap(); + let published = client + .post(format!("{base_url}/agents/{agent_id}/publish")) + .json(&json!({ "version": 1 })) + .send() + .await + .unwrap() + .json::() + .await + .unwrap(); + + assert_eq!( + bindings["bindings"][0]["tool_name"], + "crm_create_lead_agent" + ); + assert_eq!(published["published_version"], 1); + } + #[tokio::test] async fn creates_publishes_and_tests_graphql_operation() { let registry = test_registry().await; diff --git a/apps/admin-api/src/error.rs b/apps/admin-api/src/error.rs index c0bc214..644f0e0 100644 --- a/apps/admin-api/src/error.rs +++ b/apps/admin-api/src/error.rs @@ -99,6 +99,15 @@ impl From for ApiError { RegistryError::WorkspaceNotFound { workspace_id } => { Self::not_found(format!("workspace {workspace_id} was not found")) } + RegistryError::AgentNotFound { agent_id } => { + Self::not_found(format!("agent {agent_id} was not found")) + } + RegistryError::PublishedAgentNotFound { + workspace_slug, + agent_slug, + } => Self::not_found(format!( + "published agent {workspace_slug}/{agent_slug} was not found" + )), RegistryError::OperationNotFound { operation_id } => { Self::not_found(format!("operation {operation_id} was not found")) } diff --git a/apps/admin-api/src/routes.rs b/apps/admin-api/src/routes.rs index 8966bd7..8040641 100644 --- a/apps/admin-api/src/routes.rs +++ b/apps/admin-api/src/routes.rs @@ -1,3 +1,4 @@ +pub mod agents; pub mod auth_profiles; pub mod operations; pub mod workspaces; diff --git a/apps/admin-api/src/routes/agents.rs b/apps/admin-api/src/routes/agents.rs new file mode 100644 index 0000000..8d5b74c --- /dev/null +++ b/apps/admin-api/src/routes/agents.rs @@ -0,0 +1,114 @@ +use axum::{ + Json, + extract::{Path, State}, +}; +use serde::Deserialize; +use serde_json::{Value, json}; + +use crate::{ + error::ApiError, + service::{AgentBindingPayload, AgentPayload, PublishPayload}, + state::AppState, +}; + +#[derive(Deserialize)] +pub struct WorkspacePath { + pub workspace_id: String, +} + +#[derive(Deserialize)] +pub struct WorkspaceAgentPath { + pub workspace_id: String, + pub agent_id: String, +} + +#[derive(Deserialize)] +pub struct WorkspaceAgentVersionPath { + pub workspace_id: String, + pub agent_id: String, + pub version: u32, +} + +pub async fn list_agents( + Path(path): Path, + State(state): State, +) -> Result, ApiError> { + let items = state + .service + .list_agents(&path.workspace_id.as_str().into()) + .await?; + Ok(Json(json!({ "items": items }))) +} + +pub async fn create_agent( + Path(path): Path, + State(state): State, + Json(payload): Json, +) -> Result, ApiError> { + let created = state + .service + .create_agent(&path.workspace_id.as_str().into(), payload) + .await?; + Ok(Json(json!(created))) +} + +pub async fn get_agent( + Path(path): Path, + State(state): State, +) -> Result, ApiError> { + let agent = state + .service + .get_agent( + &path.workspace_id.as_str().into(), + &path.agent_id.as_str().into(), + ) + .await?; + Ok(Json(json!(agent))) +} + +pub async fn get_agent_version( + Path(path): Path, + State(state): State, +) -> Result, ApiError> { + let version = state + .service + .get_agent_version( + &path.workspace_id.as_str().into(), + &path.agent_id.as_str().into(), + path.version, + ) + .await?; + Ok(Json(json!(version))) +} + +pub async fn save_agent_bindings( + Path(path): Path, + State(state): State, + Json(payload): Json>, +) -> Result, ApiError> { + let record = state + .service + .save_agent_bindings( + &path.workspace_id.as_str().into(), + &path.agent_id.as_str().into(), + payload, + ) + .await?; + Ok(Json(json!(record))) +} + +pub async fn publish_agent( + Path(path): Path, + State(state): State, + Json(payload): Json, +) -> Result, ApiError> { + let published = state + .service + .publish_agent( + &path.workspace_id.as_str().into(), + &path.agent_id.as_str().into(), + payload.version, + ) + .await?; + Ok(Json(json!(published))) +} diff --git a/apps/admin-api/src/service.rs b/apps/admin-api/src/service.rs index 84e0ea3..68dcd0b 100644 --- a/apps/admin-api/src/service.rs +++ b/apps/admin-api/src/service.rs @@ -1,17 +1,19 @@ use std::path::PathBuf; use crank_core::{ - AuthConfig, AuthKind, AuthProfile, AuthProfileId, ConfigExport, ExportMode, GeneratedDraft, - GeneratedDraftStatus, OperationId, OperationStatus, Protocol, SampleId, Samples, Target, - Workspace, WorkspaceId, WorkspaceStatus, + Agent, AgentId, AgentOperationBinding, AgentStatus, AgentVersion, AuthConfig, AuthKind, + AuthProfile, AuthProfileId, ConfigExport, ExportMode, GeneratedDraft, GeneratedDraftStatus, + OperationId, OperationStatus, Protocol, SampleId, Samples, Target, Workspace, WorkspaceId, + WorkspaceStatus, }; use crank_mapping::{JsonPathRoot, MappingSet, infer_mapping_from_samples}; use crank_proto::{ProtoService, services_from_descriptor_set_bytes}; use crank_registry::{ - CreateVersionRequest, CreateWorkspaceRequest, OperationSampleMetadata, OperationSummary, - OperationVersionRecord, PostgresRegistry, PublishRequest, RegistryOperation, SampleKind, - SaveAuthProfileRequest, SaveDescriptorMetadataRequest, SaveSampleMetadataRequest, - UpdateWorkspaceRequest, WorkspaceRecord, + AgentSummary, AgentVersionRecord, CreateAgentRequest, CreateVersionRequest, + CreateWorkspaceRequest, OperationSampleMetadata, OperationSummary, OperationVersionRecord, + PostgresRegistry, PublishAgentRequest, PublishRequest, RegistryOperation, SampleKind, + SaveAgentBindingsRequest, SaveAuthProfileRequest, SaveDescriptorMetadataRequest, + SaveSampleMetadataRequest, UpdateWorkspaceRequest, WorkspaceRecord, }; use crank_runtime::{PreparedRequest, RuntimeError, RuntimeExecutor, RuntimeOperation}; use crank_schema::Schema; @@ -94,6 +96,44 @@ pub struct UpdateWorkspacePayload { pub settings: Option, } +#[derive(Clone, Debug, Deserialize)] +pub struct AgentPayload { + pub slug: String, + pub display_name: String, + pub description: String, + #[serde(default)] + pub instructions: Value, + #[serde(default)] + pub tool_selection_policy: Value, +} + +#[derive(Clone, Debug, Deserialize)] +pub struct AgentBindingPayload { + pub operation_id: String, + pub operation_version: u32, + pub tool_name: String, + pub tool_title: String, + pub tool_description_override: Option, + #[serde(default = "default_enabled")] + pub enabled: bool, +} + +#[derive(Clone, Debug, Serialize)] +pub struct CreatedAgentResponse { + pub agent_id: String, + pub workspace_id: String, + pub version: u32, + pub status: AgentStatus, +} + +#[derive(Clone, Debug, Serialize)] +pub struct PublishAgentResponse { + pub agent_id: String, + pub workspace_id: String, + pub published_version: u32, + pub published_at: String, +} + #[derive(Clone, Debug, Deserialize)] pub struct GenerateDraftPayload { #[serde(default)] @@ -493,6 +533,173 @@ impl AdminService { Ok(self.registry.list_auth_profiles(workspace_id).await?) } + #[instrument(skip(self))] + pub async fn list_agents( + &self, + workspace_id: &WorkspaceId, + ) -> Result, ApiError> { + self.ensure_workspace_exists(workspace_id).await?; + Ok(self.registry.list_agents(workspace_id).await?) + } + + #[instrument(skip(self))] + pub async fn get_agent( + &self, + workspace_id: &WorkspaceId, + agent_id: &AgentId, + ) -> Result { + self.registry + .get_agent_summary(workspace_id, agent_id) + .await? + .ok_or_else(|| { + ApiError::not_found(format!("agent {} was not found", agent_id.as_str())) + }) + } + + #[instrument(skip(self))] + pub async fn get_agent_version( + &self, + workspace_id: &WorkspaceId, + agent_id: &AgentId, + version: u32, + ) -> Result { + self.registry + .get_agent_version(workspace_id, agent_id, version) + .await? + .ok_or_else(|| { + ApiError::not_found(format!( + "agent version {version} for {} was not found", + agent_id.as_str() + )) + }) + } + + #[instrument(skip(self, payload), fields(workspace_id = %workspace_id.as_str(), agent_slug = %payload.slug))] + pub async fn create_agent( + &self, + workspace_id: &WorkspaceId, + payload: AgentPayload, + ) -> Result { + self.ensure_workspace_exists(workspace_id).await?; + + if self + .find_agent_by_slug(workspace_id, &payload.slug) + .await? + .is_some() + { + return Err(ApiError::conflict(format!( + "agent with slug {} already exists", + payload.slug + ))); + } + + let now = now_string()?; + let agent_id = AgentId::new(new_prefixed_id("agent")); + let agent = Agent { + id: agent_id.clone(), + workspace_id: workspace_id.clone(), + slug: payload.slug, + display_name: payload.display_name, + description: payload.description, + status: AgentStatus::Draft, + current_draft_version: 1, + latest_published_version: None, + created_at: now.clone(), + updated_at: now.clone(), + published_at: None, + }; + let version = AgentVersion { + agent_id: agent_id.clone(), + version: 1, + status: AgentStatus::Draft, + instructions: payload.instructions, + tool_selection_policy: payload.tool_selection_policy, + created_at: now, + }; + + self.registry + .create_agent(CreateAgentRequest { + agent: &agent, + version: &version, + bindings: &[], + }) + .await?; + info!(agent_id = %agent_id.as_str(), version = 1, "agent created"); + + Ok(CreatedAgentResponse { + agent_id: agent_id.as_str().to_owned(), + workspace_id: workspace_id.as_str().to_owned(), + version: 1, + status: AgentStatus::Draft, + }) + } + + #[instrument(skip(self, payload), fields(workspace_id = %workspace_id.as_str(), agent_id = %agent_id.as_str()))] + pub async fn save_agent_bindings( + &self, + workspace_id: &WorkspaceId, + agent_id: &AgentId, + payload: Vec, + ) -> Result { + let agent = self.get_agent(workspace_id, agent_id).await?; + let bindings = payload + .into_iter() + .map(|binding| AgentOperationBinding { + operation_id: OperationId::new(binding.operation_id), + operation_version: binding.operation_version, + tool_name: binding.tool_name, + tool_title: binding.tool_title, + tool_description_override: binding.tool_description_override, + enabled: binding.enabled, + }) + .collect::>(); + + self.registry + .save_agent_bindings(SaveAgentBindingsRequest { + workspace_id, + agent_id, + agent_version: agent.current_draft_version, + bindings: &bindings, + }) + .await?; + info!( + agent_id = %agent_id.as_str(), + version = agent.current_draft_version, + binding_count = bindings.len(), + "agent bindings saved" + ); + + self.get_agent_version(workspace_id, agent_id, agent.current_draft_version) + .await + } + + #[instrument(skip(self), fields(workspace_id = %workspace_id.as_str(), agent_id = %agent_id.as_str(), version))] + pub async fn publish_agent( + &self, + workspace_id: &WorkspaceId, + agent_id: &AgentId, + version: u32, + ) -> Result { + let published_at = now_string()?; + self.registry + .publish_agent(PublishAgentRequest { + workspace_id, + agent_id, + version, + published_at: &published_at, + published_by: None, + }) + .await?; + info!(agent_id = %agent_id.as_str(), version, "agent published"); + + Ok(PublishAgentResponse { + agent_id: agent_id.as_str().to_owned(), + workspace_id: workspace_id.as_str().to_owned(), + published_version: version, + published_at, + }) + } + #[instrument(skip(self, payload), fields(operation_id = %operation_id.as_str(), source_name = source_name.unwrap_or("descriptor-set.bin")))] pub async fn upload_descriptor_set( &self, @@ -940,6 +1147,19 @@ impl AdminService { .find(|operation| operation.name == name)) } + async fn find_agent_by_slug( + &self, + workspace_id: &WorkspaceId, + slug: &str, + ) -> Result, ApiError> { + Ok(self + .registry + .list_agents(workspace_id) + .await? + .into_iter() + .find(|agent| agent.slug == slug)) + } + async fn ensure_workspace_exists(&self, workspace_id: &WorkspaceId) -> Result<(), ApiError> { self.get_workspace(workspace_id).await.map(|_| ()) } @@ -1012,6 +1232,10 @@ fn default_export_mode() -> ExportMode { ExportMode::Portable } +fn default_enabled() -> bool { + true +} + fn new_prefixed_id(prefix: &str) -> String { format!("{prefix}_{}", Uuid::now_v7().simple()) } diff --git a/apps/mcp-server/src/app.rs b/apps/mcp-server/src/app.rs index 00bec4e..478f1c2 100644 --- a/apps/mcp-server/src/app.rs +++ b/apps/mcp-server/src/app.rs @@ -2,7 +2,7 @@ use std::{sync::Arc, time::Duration}; use axum::{ Json, Router, - extract::State, + extract::{Path, State}, http::{ HeaderMap, HeaderValue, StatusCode, header::{self, ACCEPT}, @@ -10,7 +10,7 @@ use axum::{ response::{IntoResponse, Response}, routing::get, }; -use crank_registry::{PostgresRegistry, RegistryOperation}; +use crank_registry::{PostgresRegistry, PublishedAgentTool}; use crank_runtime::{RuntimeError, RuntimeExecutor, RuntimeOperation}; use serde::{Deserialize, Serialize}; use serde_json::{Value, json}; @@ -54,6 +54,12 @@ struct ToolCallParams { arguments: Value, } +#[derive(Clone, Debug, Deserialize)] +struct AgentRoutePath { + workspace_slug: String, + agent_slug: String, +} + pub fn build_app( registry: PostgresRegistry, refresh_interval: Duration, @@ -68,8 +74,10 @@ pub fn build_app( Router::new() .route("/health", get(health)) - .route("/", get(mcp_get).post(mcp_post).delete(mcp_delete)) - .route("/mcp", get(mcp_get).post(mcp_post).delete(mcp_delete)) + .route( + "/v1/{workspace_slug}/{agent_slug}", + get(mcp_get).post(mcp_post).delete(mcp_delete), + ) .with_state(state) } @@ -80,25 +88,37 @@ async fn health() -> Json { })) } -async fn mcp_get() -> Response { +async fn mcp_get(Path(_path): Path) -> Response { StatusCode::METHOD_NOT_ALLOWED.into_response() } -async fn mcp_delete(State(state): State>, headers: HeaderMap) -> Response { +async fn mcp_delete( + Path(path): Path, + State(state): State>, + headers: HeaderMap, +) -> Response { match session_id_from_headers(&headers) { - Ok(Some(session_id)) => { - if state.sessions.delete(&session_id).await { - StatusCode::NO_CONTENT.into_response() - } else { - StatusCode::NOT_FOUND.into_response() + Ok(Some(session_id)) => match state.sessions.get(&session_id).await { + Some(session) + if session.workspace_slug == path.workspace_slug + && session.agent_slug == path.agent_slug => + { + if state.sessions.delete(&session_id).await { + StatusCode::NO_CONTENT.into_response() + } else { + StatusCode::NOT_FOUND.into_response() + } } - } + Some(_) => StatusCode::NOT_FOUND.into_response(), + None => StatusCode::NOT_FOUND.into_response(), + }, Ok(None) => StatusCode::BAD_REQUEST.into_response(), Err(status) => status.into_response(), } } async fn mcp_post( + Path(path): Path, State(state): State>, headers: HeaderMap, Json(message): Json, @@ -121,12 +141,15 @@ async fn mcp_post( }; match method_name(&message) { - Some("initialize") if is_request(&message) => handle_initialize(state, &message).await, + Some("initialize") if is_request(&message) => { + handle_initialize(state, &path, &message).await + } Some("notifications/initialized") if is_notification(&message) => { - handle_initialized_notification(state, &headers).await + handle_initialized_notification(state, &path, &headers).await } Some("ping") if is_request(&message) => { - let session = match require_initialized_session(&state, &headers, &message).await { + let session = match require_initialized_session(&state, &path, &headers, &message).await + { Ok(session) => session, Err(response) => return response, }; @@ -142,12 +165,17 @@ async fn mcp_post( ) } Some("tools/list") if is_request(&message) => { - let session = match require_initialized_session(&state, &headers, &message).await { + let session = match require_initialized_session(&state, &path, &headers, &message).await + { Ok(session) => session, Err(response) => return response, }; - match state.catalog.list_tools().await { + match state + .catalog + .list_tools(&session.workspace_slug, &session.agent_slug) + .await + { Ok(tools) => { let definitions = tools.iter().map(tool_definition).collect::>(); @@ -162,7 +190,8 @@ async fn mcp_post( } } Some("tools/call") if is_request(&message) => { - let session = match require_initialized_session(&state, &headers, &message).await { + let session = match require_initialized_session(&state, &path, &headers, &message).await + { Ok(session) => session, Err(response) => return response, }; @@ -184,11 +213,19 @@ async fn mcp_post( tool_call_params.arguments }; - match state.catalog.get_tool(&tool_call_params.name).await { - Ok(Some(operation)) => { + match state + .catalog + .get_tool( + &session.workspace_slug, + &session.agent_slug, + &tool_call_params.name, + ) + .await + { + Ok(Some(tool)) => { match state .runtime - .execute(&RuntimeOperation::from(operation), &arguments) + .execute(&runtime_operation(tool), &arguments) .await { Ok(output) => json_response( @@ -270,7 +307,11 @@ async fn mcp_post( } } -async fn handle_initialize(state: Arc, message: &Value) -> Response { +async fn handle_initialize( + state: Arc, + path: &AgentRoutePath, + message: &Value, +) -> Response { let initialize_params: InitializeParams = match serde_json::from_value(params(message)) { Ok(value) => value, Err(error) => { @@ -298,7 +339,10 @@ async fn handle_initialize(state: Arc, message: &Value) -> Response { Some(DEFAULT_PROTOCOL_VERSION), ); }; - let session_id = state.sessions.create(protocol_version).await; + let session_id = state + .sessions + .create(protocol_version, &path.workspace_slug, &path.agent_slug) + .await; json_response( StatusCode::OK, @@ -322,15 +366,26 @@ async fn handle_initialize(state: Arc, message: &Value) -> Response { ) } -async fn handle_initialized_notification(state: Arc, headers: &HeaderMap) -> Response { +async fn handle_initialized_notification( + state: Arc, + path: &AgentRoutePath, + headers: &HeaderMap, +) -> Response { match session_id_from_headers(headers) { - Ok(Some(session_id)) => { - if state.sessions.mark_initialized(&session_id).await { - StatusCode::ACCEPTED.into_response() - } else { - StatusCode::NOT_FOUND.into_response() + Ok(Some(session_id)) => match state.sessions.get(&session_id).await { + Some(session) + if session.workspace_slug == path.workspace_slug + && session.agent_slug == path.agent_slug => + { + if state.sessions.mark_initialized(&session_id).await { + StatusCode::ACCEPTED.into_response() + } else { + StatusCode::NOT_FOUND.into_response() + } } - } + Some(_) => StatusCode::NOT_FOUND.into_response(), + None => StatusCode::NOT_FOUND.into_response(), + }, Ok(None) => StatusCode::BAD_REQUEST.into_response(), Err(status) => status.into_response(), } @@ -338,6 +393,7 @@ async fn handle_initialized_notification(state: Arc, headers: &HeaderM async fn require_initialized_session( state: &Arc, + path: &AgentRoutePath, headers: &HeaderMap, message: &Value, ) -> Result { @@ -351,6 +407,10 @@ async fn require_initialized_session( return Err(StatusCode::NOT_FOUND.into_response()); }; + if session.workspace_slug != path.workspace_slug || session.agent_slug != path.agent_slug { + return Err(StatusCode::NOT_FOUND.into_response()); + } + if !session.initialized { return Err(json_response( StatusCode::OK, @@ -471,15 +531,23 @@ fn json_response( response } -fn tool_definition(operation: &RegistryOperation) -> Value { +fn tool_definition(tool: &PublishedAgentTool) -> Value { json!({ - "name": operation.name, - "title": operation.tool_description.title, - "description": operation.tool_description.description, - "inputSchema": schema_to_json_schema(&operation.input_schema) + "name": tool.tool_name, + "title": tool.tool_title, + "description": tool.tool_description, + "inputSchema": schema_to_json_schema(&tool.operation.input_schema) }) } +fn runtime_operation(tool: PublishedAgentTool) -> RuntimeOperation { + let mut operation = RuntimeOperation::from(tool.operation); + operation.tool_name = tool.tool_name; + operation.tool_description.title = tool.tool_title; + operation.tool_description.description = tool.tool_description; + operation +} + fn schema_to_json_schema(schema: &crank_schema::Schema) -> Value { match schema.kind { crank_schema::SchemaKind::Object => { diff --git a/apps/mcp-server/src/catalog.rs b/apps/mcp-server/src/catalog.rs index 49e19a9..3a2df58 100644 --- a/apps/mcp-server/src/catalog.rs +++ b/apps/mcp-server/src/catalog.rs @@ -4,7 +4,7 @@ use std::{ time::{Duration, Instant}, }; -use crank_registry::{PostgresRegistry, RegistryError, RegistryOperation}; +use crank_registry::{PostgresRegistry, PublishedAgentTool, RegistryError}; use tokio::sync::RwLock; use tracing::info; @@ -12,14 +12,20 @@ use tracing::info; pub struct PublishedToolCatalog { registry: PostgresRegistry, refresh_interval: Duration, - cached: Arc>, + cached: Arc>>, +} + +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +struct CatalogKey { + workspace_slug: String, + agent_slug: String, } #[derive(Default)] struct CachedCatalog { loaded_at: Option, - tools: Vec, - tools_by_name: HashMap, + tools: Vec, + tools_by_name: HashMap, } impl PublishedToolCatalog { @@ -27,30 +33,47 @@ impl PublishedToolCatalog { Self { registry, refresh_interval, - cached: Arc::new(RwLock::new(CachedCatalog::default())), + cached: Arc::new(RwLock::new(HashMap::new())), } } - pub async fn list_tools(&self) -> Result, RegistryError> { - self.refresh_if_stale().await?; + pub async fn list_tools( + &self, + workspace_slug: &str, + agent_slug: &str, + ) -> Result, RegistryError> { + self.refresh_if_stale(workspace_slug, agent_slug).await?; let guard = self.cached.read().await; - Ok(guard.tools.clone()) + Ok(guard + .get(&CatalogKey::new(workspace_slug, agent_slug)) + .map(|entry| entry.tools.clone()) + .unwrap_or_default()) } pub async fn get_tool( &self, + workspace_slug: &str, + agent_slug: &str, tool_name: &str, - ) -> Result, RegistryError> { - self.refresh_if_stale().await?; + ) -> Result, RegistryError> { + self.refresh_if_stale(workspace_slug, agent_slug).await?; let guard = self.cached.read().await; - Ok(guard.tools_by_name.get(tool_name).cloned()) + Ok(guard + .get(&CatalogKey::new(workspace_slug, agent_slug)) + .and_then(|entry| entry.tools_by_name.get(tool_name)) + .cloned()) } - async fn refresh_if_stale(&self) -> Result<(), RegistryError> { + async fn refresh_if_stale( + &self, + workspace_slug: &str, + agent_slug: &str, + ) -> Result<(), RegistryError> { + let key = CatalogKey::new(workspace_slug, agent_slug); let should_refresh = { let guard = self.cached.read().await; - match guard.loaded_at { + match guard.get(&key).and_then(|entry| entry.loaded_at) { Some(loaded_at) => loaded_at.elapsed() >= self.refresh_interval, None => true, } @@ -60,25 +83,55 @@ impl PublishedToolCatalog { return Ok(()); } - let tools = self.registry.list_published_operations().await?; + let tools = match self + .registry + .get_published_agent_tools_by_slug(workspace_slug, agent_slug) + .await + { + Ok(tools) => tools, + Err(RegistryError::PublishedAgentNotFound { .. }) => Vec::new(), + Err(error) => return Err(error), + }; let tools_by_name = tools .iter() .cloned() - .map(|operation| (operation.name.clone(), operation)) + .map(|tool| (tool.tool_name.clone(), tool)) .collect::>(); let mut guard = self.cached.write().await; - let previous_count = guard.tools.len(); + let previous_count = guard + .get(&key) + .map(|entry| entry.tools.len()) + .unwrap_or_default(); - guard.loaded_at = Some(Instant::now()); - guard.tools = tools; - guard.tools_by_name = tools_by_name; + guard.insert( + key, + CachedCatalog { + loaded_at: Some(Instant::now()), + tools, + tools_by_name, + }, + ); info!( - published_tool_count = guard.tools.len(), + workspace_slug, + agent_slug, + published_tool_count = guard + .get(&CatalogKey::new(workspace_slug, agent_slug)) + .map(|entry| entry.tools.len()) + .unwrap_or_default(), previous_published_tool_count = previous_count, - "published tool catalog refreshed" + "published agent catalog refreshed" ); Ok(()) } } + +impl CatalogKey { + fn new(workspace_slug: &str, agent_slug: &str) -> Self { + Self { + workspace_slug: workspace_slug.to_owned(), + agent_slug: agent_slug.to_owned(), + } + } +} diff --git a/apps/mcp-server/src/main.rs b/apps/mcp-server/src/main.rs index 5217e49..8ccaede 100644 --- a/apps/mcp-server/src/main.rs +++ b/apps/mcp-server/src/main.rs @@ -51,12 +51,14 @@ mod tests { use axum::{Json, Router, http::header, routing::post}; use crank_adapter_grpc::test_support as grpc_test_support; use crank_core::{ - DescriptorId, ExecutionConfig, GraphqlOperationType, GraphqlTarget, GrpcTarget, HttpMethod, - Operation, OperationId, OperationStatus, Protocol, RestTarget, Target, ToolDescription, - WorkspaceId, + Agent, AgentId, AgentOperationBinding, AgentStatus, AgentVersion, DescriptorId, + ExecutionConfig, GraphqlOperationType, GraphqlTarget, GrpcTarget, HttpMethod, Operation, + OperationId, OperationStatus, Protocol, RestTarget, Target, ToolDescription, WorkspaceId, }; use crank_mapping::{MappingRule, MappingSet}; - use crank_registry::{PostgresRegistry, PublishRequest}; + use crank_registry::{ + CreateAgentRequest, PostgresRegistry, PublishAgentRequest, PublishRequest, + }; use crank_schema::{Schema, SchemaKind}; use serde_json::{Value, json}; use sqlx::{Executor, postgres::PgPoolOptions}; @@ -68,6 +70,10 @@ mod tests { WorkspaceId::new("ws_default") } + fn test_workspace_slug() -> &'static str { + "default" + } + #[tokio::test] async fn initializes_lists_and_calls_published_tool_via_mcp() { let registry = test_registry().await; @@ -88,6 +94,7 @@ mod tests { }) .await .unwrap(); + publish_agent_for_operation(®istry, &operation, "sales-rest").await; let base_url = spawn_mcp_server(build_app( registry, @@ -96,11 +103,12 @@ mod tests { )) .await; let client = reqwest::Client::new(); - let initialized_session = initialize_session(&client, &base_url).await; + let mcp_url = agent_mcp_url(&base_url, "sales-rest"); + let initialized_session = initialize_session(&client, &mcp_url).await; let tools = post_jsonrpc( &client, - &base_url, + &mcp_url, Some(&initialized_session), json!({ "jsonrpc": "2.0", @@ -112,7 +120,7 @@ mod tests { .await; let call_result = post_jsonrpc( &client, - &base_url, + &mcp_url, Some(&initialized_session), json!({ "jsonrpc": "2.0", @@ -156,6 +164,7 @@ mod tests { }) .await .unwrap(); + publish_agent_for_operation(®istry, &operation, "sales-graphql").await; let base_url = spawn_mcp_server(build_app( registry, @@ -164,11 +173,12 @@ mod tests { )) .await; let client = reqwest::Client::new(); - let initialized_session = initialize_session(&client, &base_url).await; + let mcp_url = agent_mcp_url(&base_url, "sales-graphql"); + let initialized_session = initialize_session(&client, &mcp_url).await; let call_result = post_jsonrpc( &client, - &base_url, + &mcp_url, Some(&initialized_session), json!({ "jsonrpc": "2.0", @@ -211,6 +221,7 @@ mod tests { }) .await .unwrap(); + publish_agent_for_operation(®istry, &operation, "sales-grpc").await; let base_url = spawn_mcp_server(build_app( registry, @@ -219,11 +230,12 @@ mod tests { )) .await; let client = reqwest::Client::new(); - let initialized_session = initialize_session(&client, &base_url).await; + let mcp_url = agent_mcp_url(&base_url, "sales-grpc"); + let initialized_session = initialize_session(&client, &mcp_url).await; let call_result = post_jsonrpc( &client, - &base_url, + &mcp_url, Some(&initialized_session), json!({ "jsonrpc": "2.0", @@ -256,8 +268,9 @@ mod tests { )) .await; let client = reqwest::Client::new(); + let mcp_url = agent_mcp_url(&base_url, "sales-init"); let initialize_response = client - .post(format!("{base_url}/mcp")) + .post(&mcp_url) .header(header::ACCEPT, "application/json, text/event-stream") .json(&json!({ "jsonrpc": "2.0", @@ -278,7 +291,7 @@ mod tests { .unwrap() .to_owned(); let tools_list = client - .post(format!("{base_url}/mcp")) + .post(&mcp_url) .header(header::ACCEPT, "application/json, text/event-stream") .header("MCP-Session-Id", session_id) .header("MCP-Protocol-Version", "2025-11-25") @@ -309,11 +322,12 @@ mod tests { )) .await; let client = reqwest::Client::new(); - let initialized_session = initialize_session(&client, &base_url).await; + let mcp_url = agent_mcp_url(&base_url, "sales-refresh"); + let initialized_session = initialize_session(&client, &mcp_url).await; let before_publish = post_jsonrpc( &client, - &base_url, + &mcp_url, Some(&initialized_session), json!({ "jsonrpc": "2.0", @@ -339,10 +353,11 @@ mod tests { }) .await .unwrap(); + publish_agent_for_operation(®istry, &operation, "sales-refresh").await; let after_publish = post_jsonrpc( &client, - &base_url, + &mcp_url, Some(&initialized_session), json!({ "jsonrpc": "2.0", @@ -360,9 +375,87 @@ mod tests { ); } - async fn initialize_session(client: &reqwest::Client, base_url: &str) -> String { + #[tokio::test] + async fn lists_only_bound_tools_for_agent_context() { + let registry = test_registry().await; + let upstream_base_url = spawn_upstream_server().await; + let operation_a = test_operation(&upstream_base_url, "crm_create_lead"); + let operation_b = test_operation(&upstream_base_url, "crm_update_lead"); + + for operation in [&operation_a, &operation_b] { + registry + .create_operation(&test_workspace_id(), operation, Some("alice")) + .await + .unwrap(); + registry + .publish_operation(PublishRequest { + workspace_id: &test_workspace_id(), + operation_id: &operation.id, + version: 1, + published_at: "2026-03-26T10:00:00Z", + published_by: Some("alice"), + }) + .await + .unwrap(); + } + + publish_agent_with_bindings( + ®istry, + "sales-a", + vec![binding_for_operation(&operation_a)], + ) + .await; + publish_agent_with_bindings( + ®istry, + "sales-b", + vec![binding_for_operation(&operation_b)], + ) + .await; + + let base_url = spawn_mcp_server(build_app( + registry, + Duration::from_millis(0), + Some("https://crank.example.com".to_owned()), + )) + .await; + let client = reqwest::Client::new(); + let agent_a_url = agent_mcp_url(&base_url, "sales-a"); + let agent_b_url = agent_mcp_url(&base_url, "sales-b"); + let session_a = initialize_session(&client, &agent_a_url).await; + let session_b = initialize_session(&client, &agent_b_url).await; + + let tools_a = post_jsonrpc( + &client, + &agent_a_url, + Some(&session_a), + json!({ + "jsonrpc": "2.0", + "id": 2, + "method": "tools/list", + "params": {} + }), + ) + .await; + let tools_b = post_jsonrpc( + &client, + &agent_b_url, + Some(&session_b), + json!({ + "jsonrpc": "2.0", + "id": 2, + "method": "tools/list", + "params": {} + }), + ) + .await; + + assert_eq!(tools_a["result"]["tools"][0]["name"], "crm_create_lead"); + assert_eq!(tools_b["result"]["tools"][0]["name"], "crm_update_lead"); + } + + async fn initialize_session(client: &reqwest::Client, mcp_url: &str) -> String { let initialize_response = client - .post(format!("{base_url}/mcp")) + .post(mcp_url) .header(header::ACCEPT, "application/json, text/event-stream") .json(&json!({ "jsonrpc": "2.0", @@ -384,7 +477,7 @@ mod tests { .to_owned(); let initialized_response = client - .post(format!("{base_url}/mcp")) + .post(mcp_url) .header(header::ACCEPT, "application/json, text/event-stream") .header("MCP-Session-Id", &session_id) .header("MCP-Protocol-Version", "2025-11-25") @@ -404,12 +497,12 @@ mod tests { async fn post_jsonrpc( client: &reqwest::Client, - base_url: &str, + mcp_url: &str, session_id: Option<&str>, payload: Value, ) -> Value { let mut request = client - .post(format!("{base_url}/mcp")) + .post(mcp_url) .header(header::ACCEPT, "application/json, text/event-stream") .header("MCP-Protocol-Version", "2025-11-25"); @@ -462,6 +555,78 @@ mod tests { format!("http://{}", address) } + fn agent_mcp_url(base_url: &str, agent_slug: &str) -> String { + format!("{base_url}/v1/{}/{}", test_workspace_slug(), agent_slug) + } + + async fn publish_agent_for_operation( + registry: &PostgresRegistry, + operation: &Operation, + agent_slug: &str, + ) { + publish_agent_with_bindings(registry, agent_slug, vec![binding_for_operation(operation)]) + .await; + } + + async fn publish_agent_with_bindings( + registry: &PostgresRegistry, + agent_slug: &str, + bindings: Vec, + ) { + let agent_id = AgentId::new(format!("agent_{agent_slug}")); + let agent = Agent { + id: agent_id.clone(), + workspace_id: test_workspace_id(), + slug: agent_slug.to_owned(), + display_name: format!("Agent {agent_slug}"), + description: "Curated MCP toolset".to_owned(), + status: AgentStatus::Draft, + current_draft_version: 1, + latest_published_version: None, + created_at: "2026-03-26T10:00:00Z".to_owned(), + updated_at: "2026-03-26T10:00:00Z".to_owned(), + published_at: None, + }; + let version = AgentVersion { + agent_id: agent_id.clone(), + version: 1, + status: AgentStatus::Draft, + instructions: json!({}), + tool_selection_policy: json!({}), + created_at: "2026-03-26T10:00:00Z".to_owned(), + }; + + registry + .create_agent(CreateAgentRequest { + agent: &agent, + version: &version, + bindings: &bindings, + }) + .await + .unwrap(); + registry + .publish_agent(PublishAgentRequest { + workspace_id: &test_workspace_id(), + agent_id: &agent_id, + version: 1, + published_at: "2026-03-26T10:00:00Z", + published_by: Some("alice"), + }) + .await + .unwrap(); + } + + fn binding_for_operation(operation: &Operation) -> AgentOperationBinding { + AgentOperationBinding { + operation_id: operation.id.clone(), + operation_version: 1, + tool_name: operation.name.clone(), + tool_title: operation.tool_description.title.clone(), + tool_description_override: None, + enabled: true, + } + } + async fn create_lead(Json(payload): Json) -> Json { Json(json!({ "id": "lead_123", diff --git a/apps/mcp-server/src/session.rs b/apps/mcp-server/src/session.rs index 5335ebd..4f55a37 100644 --- a/apps/mcp-server/src/session.rs +++ b/apps/mcp-server/src/session.rs @@ -12,6 +12,8 @@ pub struct SessionStore { pub struct SessionState { pub protocol_version: String, pub initialized: bool, + pub workspace_slug: String, + pub agent_slug: String, } impl SessionStore { @@ -21,7 +23,12 @@ impl SessionStore { } } - pub async fn create(&self, protocol_version: &str) -> String { + pub async fn create( + &self, + protocol_version: &str, + workspace_slug: &str, + agent_slug: &str, + ) -> String { let session_id = Uuid::now_v7().to_string(); let mut guard = self.inner.write().await; @@ -30,6 +37,8 @@ impl SessionStore { SessionState { protocol_version: protocol_version.to_owned(), initialized: false, + workspace_slug: workspace_slug.to_owned(), + agent_slug: agent_slug.to_owned(), }, ); diff --git a/crates/crank-core/src/agent.rs b/crates/crank-core/src/agent.rs new file mode 100644 index 0000000..2fbbd00 --- /dev/null +++ b/crates/crank-core/src/agent.rs @@ -0,0 +1,47 @@ +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +use crate::ids::{AgentId, OperationId, WorkspaceId}; + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum AgentStatus { + Draft, + Published, + Archived, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct Agent { + pub id: AgentId, + pub workspace_id: WorkspaceId, + pub slug: String, + pub display_name: String, + pub description: String, + pub status: AgentStatus, + pub current_draft_version: u32, + pub latest_published_version: Option, + pub created_at: String, + pub updated_at: String, + pub published_at: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct AgentVersion { + pub agent_id: AgentId, + pub version: u32, + pub status: AgentStatus, + pub instructions: Value, + pub tool_selection_policy: Value, + pub created_at: String, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct AgentOperationBinding { + pub operation_id: OperationId, + pub operation_version: u32, + pub tool_name: String, + pub tool_title: String, + pub tool_description_override: Option, + pub enabled: bool, +} diff --git a/crates/crank-core/src/ids.rs b/crates/crank-core/src/ids.rs index a6eab93..eb281f5 100644 --- a/crates/crank-core/src/ids.rs +++ b/crates/crank-core/src/ids.rs @@ -42,3 +42,4 @@ define_id!(SampleId); define_id!(AuthProfileId); define_id!(WorkspaceId); define_id!(UserId); +define_id!(AgentId); diff --git a/crates/crank-core/src/lib.rs b/crates/crank-core/src/lib.rs index 3ac1737..56f186c 100644 --- a/crates/crank-core/src/lib.rs +++ b/crates/crank-core/src/lib.rs @@ -1,14 +1,18 @@ +pub mod agent; pub mod auth; pub mod ids; pub mod operation; pub mod protocol; pub mod workspace; +pub use agent::{Agent, AgentOperationBinding, AgentStatus, AgentVersion}; pub use auth::{ ApiKeyHeaderAuthConfig, ApiKeyQueryAuthConfig, AuthConfig, AuthProfile, BasicAuthConfig, BearerAuthConfig, SecretRef, }; -pub use ids::{AuthProfileId, DescriptorId, OperationId, SampleId, ToolId, UserId, WorkspaceId}; +pub use ids::{ + AgentId, AuthProfileId, DescriptorId, OperationId, SampleId, ToolId, UserId, WorkspaceId, +}; pub use operation::{ ConfigExport, ExecutionConfig, GeneratedDraft, GeneratedDraftStatus, GraphqlTarget, GrpcProtocolOptions, GrpcTarget, Operation, OperationStatus, ProtocolOptions, RestTarget, diff --git a/crates/crank-registry/src/error.rs b/crates/crank-registry/src/error.rs index be68a2e..19911e1 100644 --- a/crates/crank-registry/src/error.rs +++ b/crates/crank-registry/src/error.rs @@ -10,6 +10,13 @@ pub enum RegistryError { WorkspaceNotFound { workspace_id: String }, #[error("workspace with slug {slug} already exists")] WorkspaceSlugAlreadyExists { slug: String }, + #[error("agent {agent_id} was not found")] + AgentNotFound { agent_id: String }, + #[error("published agent {workspace_slug}/{agent_slug} was not found")] + PublishedAgentNotFound { + workspace_slug: String, + agent_slug: String, + }, #[error("operation {operation_id} already exists")] OperationAlreadyExists { operation_id: String }, #[error("operation {operation_id} was not found")] diff --git a/crates/crank-registry/src/lib.rs b/crates/crank-registry/src/lib.rs index 2758fe6..e7cc131 100644 --- a/crates/crank-registry/src/lib.rs +++ b/crates/crank-registry/src/lib.rs @@ -5,10 +5,12 @@ mod postgres; pub use error::RegistryError; pub use model::{ - CreateVersionRequest, CreateWorkspaceRequest, CreateYamlImportJobRequest, DescriptorKind, - DescriptorMetadata, OperationSampleMetadata, OperationSummary, OperationVersionRecord, - PublishRequest, RegistryOperation, SampleKind, SaveAuthProfileRequest, - SaveDescriptorMetadataRequest, SaveSampleMetadataRequest, UpdateWorkspaceRequest, - WorkspaceRecord, YamlImportJob, YamlImportJobCompletion, YamlImportJobId, YamlImportJobStatus, + AgentSummary, AgentVersionRecord, CreateAgentRequest, CreateVersionRequest, + CreateWorkspaceRequest, CreateYamlImportJobRequest, DescriptorKind, DescriptorMetadata, + OperationSampleMetadata, OperationSummary, OperationVersionRecord, PublishAgentRequest, + PublishRequest, PublishedAgentTool, RegistryOperation, SampleKind, SaveAgentBindingsRequest, + SaveAuthProfileRequest, SaveDescriptorMetadataRequest, SaveSampleMetadataRequest, + UpdateWorkspaceRequest, WorkspaceRecord, YamlImportJob, YamlImportJobCompletion, + YamlImportJobId, YamlImportJobStatus, }; pub use postgres::PostgresRegistry; diff --git a/crates/crank-registry/src/migrations.rs b/crates/crank-registry/src/migrations.rs index c3ec03a..08ce2d3 100644 --- a/crates/crank-registry/src/migrations.rs +++ b/crates/crank-registry/src/migrations.rs @@ -191,5 +191,78 @@ pub async fn apply_postgres(pool: &PgPool) -> Result<(), sqlx::Error> { .execute(pool) .await?; + query( + "create table if not exists agents ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + slug text not null, + display_name text not null, + description text not null, + status text not null, + current_draft_version integer not null default 1, + latest_published_version integer null, + created_at timestamptz not null, + updated_at timestamptz not null, + published_at timestamptz null + )", + ) + .execute(pool) + .await?; + + query( + "create unique index if not exists agents_workspace_slug_idx on agents(workspace_id, slug)", + ) + .execute(pool) + .await?; + + query( + "create table if not exists agent_versions ( + agent_id text not null references agents(id) on delete cascade, + version integer not null, + status text not null, + instructions_json jsonb not null, + tool_selection_policy_json jsonb not null, + created_at timestamptz not null, + primary key (agent_id, version) + )", + ) + .execute(pool) + .await?; + + query( + "create table if not exists agent_operation_bindings ( + agent_id text not null references agents(id) on delete cascade, + agent_version integer not null, + operation_id text not null references operations(id) on delete cascade, + operation_version integer not null, + tool_name text not null, + tool_title text not null, + tool_description_override text null, + enabled boolean not null default true, + foreign key (agent_id, agent_version) references agent_versions(agent_id, version) on delete cascade, + foreign key (operation_id, operation_version) references operation_versions(operation_id, version) on delete cascade + )", + ) + .execute(pool) + .await?; + + query( + "create unique index if not exists agent_bindings_tool_name_idx on agent_operation_bindings(agent_id, agent_version, tool_name)", + ) + .execute(pool) + .await?; + + query( + "create table if not exists published_agents ( + agent_id text primary key references agents(id) on delete cascade, + version integer not null, + published_at timestamptz not null, + published_by text null, + foreign key (agent_id, version) references agent_versions(agent_id, version) on delete cascade + )", + ) + .execute(pool) + .await?; + Ok(()) } diff --git a/crates/crank-registry/src/model.rs b/crates/crank-registry/src/model.rs index 2ee1b0d..75f27ce 100644 --- a/crates/crank-registry/src/model.rs +++ b/crates/crank-registry/src/model.rs @@ -1,6 +1,7 @@ use crank_core::{ - AuthProfile, DescriptorId, ExportMode, Operation, OperationId, OperationStatus, Protocol, - SampleId, Workspace, WorkspaceId, + Agent, AgentId, AgentOperationBinding, AgentStatus, AgentVersion, AuthProfile, DescriptorId, + ExportMode, Operation, OperationId, OperationStatus, Protocol, SampleId, Workspace, + WorkspaceId, }; use crank_mapping::MappingSet; use crank_schema::Schema; @@ -45,6 +46,44 @@ pub struct WorkspaceRecord { pub workspace: Workspace, } +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct AgentSummary { + pub id: AgentId, + pub workspace_id: WorkspaceId, + pub slug: String, + pub display_name: String, + pub description: String, + pub status: AgentStatus, + pub current_draft_version: u32, + pub latest_published_version: Option, + pub created_at: String, + pub updated_at: String, + pub published_at: Option, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct AgentVersionRecord { + pub agent_id: AgentId, + pub workspace_id: WorkspaceId, + pub version: u32, + pub status: AgentStatus, + pub created_at: String, + pub snapshot: AgentVersion, + pub bindings: Vec, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct PublishedAgentTool { + pub workspace_id: WorkspaceId, + pub workspace_slug: String, + pub agent_id: AgentId, + pub agent_slug: String, + pub operation: RegistryOperation, + pub tool_name: String, + pub tool_title: String, + pub tool_description: String, +} + #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] pub struct OperationSummary { pub id: OperationId, @@ -185,6 +224,30 @@ pub struct UpdateWorkspaceRequest<'a> { pub workspace: &'a Workspace, } +#[derive(Clone, Debug, PartialEq)] +pub struct CreateAgentRequest<'a> { + pub agent: &'a Agent, + pub version: &'a AgentVersion, + pub bindings: &'a [AgentOperationBinding], +} + +#[derive(Clone, Debug, PartialEq)] +pub struct SaveAgentBindingsRequest<'a> { + pub workspace_id: &'a WorkspaceId, + pub agent_id: &'a AgentId, + pub agent_version: u32, + pub bindings: &'a [AgentOperationBinding], +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct PublishAgentRequest<'a> { + pub workspace_id: &'a WorkspaceId, + pub agent_id: &'a AgentId, + pub version: u32, + pub published_at: &'a str, + pub published_by: Option<&'a str>, +} + #[derive(Clone, Debug, PartialEq, Eq)] pub struct SaveSampleMetadataRequest<'a> { pub sample: &'a OperationSampleMetadata, diff --git a/crates/crank-registry/src/postgres.rs b/crates/crank-registry/src/postgres.rs index fa3ddc5..6ee8a15 100644 --- a/crates/crank-registry/src/postgres.rs +++ b/crates/crank-registry/src/postgres.rs @@ -1,4 +1,7 @@ -use crank_core::{AuthProfile, OperationId, OperationStatus, Workspace, WorkspaceId}; +use crank_core::{ + AgentId, AgentOperationBinding, AgentStatus, AgentVersion, AuthProfile, OperationId, + OperationStatus, Workspace, WorkspaceId, +}; use serde::{Serialize, de::DeserializeOwned}; use serde_json::Value; use sqlx::{ @@ -11,11 +14,13 @@ use crate::{ error::RegistryError, migrations, model::{ - CreateVersionRequest, CreateWorkspaceRequest, CreateYamlImportJobRequest, - DescriptorMetadata, OperationSampleMetadata, OperationSummary, OperationVersionRecord, - PublishRequest, RegistryOperation, SaveAuthProfileRequest, SaveDescriptorMetadataRequest, - SaveSampleMetadataRequest, UpdateWorkspaceRequest, WorkspaceRecord, YamlImportJob, - YamlImportJobCompletion, YamlImportJobId, YamlImportJobStatus, + AgentSummary, AgentVersionRecord, CreateAgentRequest, CreateVersionRequest, + CreateWorkspaceRequest, CreateYamlImportJobRequest, DescriptorMetadata, + OperationSampleMetadata, OperationSummary, OperationVersionRecord, PublishAgentRequest, + PublishRequest, PublishedAgentTool, RegistryOperation, SaveAgentBindingsRequest, + SaveAuthProfileRequest, SaveDescriptorMetadataRequest, SaveSampleMetadataRequest, + UpdateWorkspaceRequest, WorkspaceRecord, YamlImportJob, YamlImportJobCompletion, + YamlImportJobId, YamlImportJobStatus, }, }; @@ -158,6 +163,312 @@ impl PostgresRegistry { } } + pub async fn list_agents( + &self, + workspace_id: &WorkspaceId, + ) -> Result, RegistryError> { + let rows = sqlx::query( + "select + id, + workspace_id, + slug, + display_name, + description, + status, + current_draft_version, + latest_published_version, + to_char(created_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as created_at, + to_char(updated_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as updated_at, + to_char(published_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as published_at + from agents + where workspace_id = $1 + order by slug asc", + ) + .bind(workspace_id.as_str()) + .fetch_all(&self.pool) + .await?; + + rows.iter().map(map_agent_summary).collect() + } + + pub async fn get_agent_summary( + &self, + workspace_id: &WorkspaceId, + agent_id: &AgentId, + ) -> Result, RegistryError> { + let row = sqlx::query( + "select + id, + workspace_id, + slug, + display_name, + description, + status, + current_draft_version, + latest_published_version, + to_char(created_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as created_at, + to_char(updated_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as updated_at, + to_char(published_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as published_at + from agents + where workspace_id = $1 and id = $2", + ) + .bind(workspace_id.as_str()) + .bind(agent_id.as_str()) + .fetch_optional(&self.pool) + .await?; + + row.as_ref().map(map_agent_summary).transpose() + } + + pub async fn create_agent(&self, request: CreateAgentRequest<'_>) -> Result<(), RegistryError> { + let mut tx = self.pool.begin().await?; + + sqlx::query( + "insert into agents ( + id, + workspace_id, + slug, + display_name, + description, + status, + current_draft_version, + latest_published_version, + created_at, + updated_at, + published_at + ) values ( + $1, $2, $3, $4, $5, $6, $7, $8, $9::timestamptz, $10::timestamptz, $11::timestamptz + )", + ) + .bind(request.agent.id.as_str()) + .bind(request.agent.workspace_id.as_str()) + .bind(&request.agent.slug) + .bind(&request.agent.display_name) + .bind(&request.agent.description) + .bind(serialize_enum_text(&request.agent.status, "status")?) + .bind(to_db_version(request.agent.current_draft_version)) + .bind(request.agent.latest_published_version.map(to_db_version)) + .bind(&request.agent.created_at) + .bind(&request.agent.updated_at) + .bind(request.agent.published_at.as_deref()) + .execute(&mut *tx) + .await?; + + insert_agent_version_row(&mut tx, request.version).await?; + replace_agent_bindings_rows( + &mut tx, + &request.agent.id, + request.version.version, + request.bindings, + ) + .await?; + + tx.commit().await?; + Ok(()) + } + + pub async fn get_agent_version( + &self, + workspace_id: &WorkspaceId, + agent_id: &AgentId, + version: u32, + ) -> Result, RegistryError> { + let summary = match self.get_agent_summary(workspace_id, agent_id).await? { + Some(value) => value, + None => return Ok(None), + }; + + let row = sqlx::query( + "select + agent_id, + version, + status, + instructions_json, + tool_selection_policy_json, + to_char(created_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as created_at + from agent_versions + where agent_id = $1 and version = $2", + ) + .bind(agent_id.as_str()) + .bind(to_db_version(version)) + .fetch_optional(&self.pool) + .await?; + + let Some(row) = row else { + return Ok(None); + }; + + let bindings = self.list_agent_bindings(agent_id, version).await?; + Ok(Some(map_agent_version_record(&summary, &row, bindings)?)) + } + + pub async fn save_agent_bindings( + &self, + request: SaveAgentBindingsRequest<'_>, + ) -> Result<(), RegistryError> { + if self + .get_agent_summary(request.workspace_id, request.agent_id) + .await? + .is_none() + { + return Err(RegistryError::AgentNotFound { + agent_id: request.agent_id.as_str().to_owned(), + }); + } + + let mut tx = self.pool.begin().await?; + replace_agent_bindings_rows( + &mut tx, + request.agent_id, + request.agent_version, + request.bindings, + ) + .await?; + tx.commit().await?; + Ok(()) + } + + pub async fn publish_agent( + &self, + request: PublishAgentRequest<'_>, + ) -> Result<(), RegistryError> { + if self + .get_agent_version(request.workspace_id, request.agent_id, request.version) + .await? + .is_none() + { + return Err(RegistryError::AgentNotFound { + agent_id: request.agent_id.as_str().to_owned(), + }); + } + + let mut tx = self.pool.begin().await?; + + sqlx::query( + "insert into published_agents ( + agent_id, + version, + published_at, + published_by + ) values ($1, $2, $3::timestamptz, $4) + on conflict(agent_id) do update set + version = excluded.version, + published_at = excluded.published_at, + published_by = excluded.published_by", + ) + .bind(request.agent_id.as_str()) + .bind(to_db_version(request.version)) + .bind(request.published_at) + .bind(request.published_by) + .execute(&mut *tx) + .await?; + + sqlx::query( + "update agent_versions + set status = $1 + where agent_id = $2 and version = $3", + ) + .bind(serialize_enum_text(&AgentStatus::Published, "status")?) + .bind(request.agent_id.as_str()) + .bind(to_db_version(request.version)) + .execute(&mut *tx) + .await?; + + sqlx::query( + "update agents + set status = $1, + latest_published_version = $2, + published_at = $3::timestamptz, + updated_at = $4::timestamptz + where id = $5 and workspace_id = $6", + ) + .bind(serialize_enum_text(&AgentStatus::Published, "status")?) + .bind(to_db_version(request.version)) + .bind(request.published_at) + .bind(request.published_at) + .bind(request.agent_id.as_str()) + .bind(request.workspace_id.as_str()) + .execute(&mut *tx) + .await?; + + tx.commit().await?; + Ok(()) + } + + pub async fn get_published_agent_tools_by_slug( + &self, + workspace_slug: &str, + agent_slug: &str, + ) -> Result, RegistryError> { + let rows = sqlx::query( + "select + w.id as workspace_id, + w.slug as workspace_slug, + a.id as agent_id, + a.slug as agent_slug, + b.tool_name, + b.tool_title, + coalesce(b.tool_description_override, ov.tool_description_json->>'description') as tool_description, + o.id, + o.name, + o.display_name, + o.protocol, + to_char(o.created_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as operation_created_at, + to_char(o.updated_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as operation_updated_at, + to_char(o.published_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as operation_published_at, + ov.version, + ov.status, + ov.target_json, + ov.input_schema_json, + ov.output_schema_json, + ov.input_mapping_json, + ov.output_mapping_json, + ov.execution_config_json, + ov.tool_description_json, + ov.samples_json, + ov.generated_draft_json, + ov.config_export_json, + ov.change_note, + to_char(ov.created_at at time zone 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') as created_at, + ov.created_by + from workspaces w + join agents a on a.workspace_id = w.id + join published_agents pa on pa.agent_id = a.id + join agent_operation_bindings b on b.agent_id = a.id and b.agent_version = pa.version + join operation_versions ov on ov.operation_id = b.operation_id and ov.version = b.operation_version + join operations o on o.id = ov.operation_id and o.workspace_id = w.id + where w.slug = $1 and a.slug = $2 and b.enabled = true + order by b.tool_name asc", + ) + .bind(workspace_slug) + .bind(agent_slug) + .fetch_all(&self.pool) + .await?; + + if rows.is_empty() { + let exists = sqlx::query( + "select 1 + from workspaces w + join agents a on a.workspace_id = w.id + join published_agents pa on pa.agent_id = a.id + where w.slug = $1 and a.slug = $2", + ) + .bind(workspace_slug) + .bind(agent_slug) + .fetch_optional(&self.pool) + .await?; + + if exists.is_none() { + return Err(RegistryError::PublishedAgentNotFound { + workspace_slug: workspace_slug.to_owned(), + agent_slug: agent_slug.to_owned(), + }); + } + } + + rows.iter().map(map_published_agent_tool).collect() + } + async fn connect_in_schema( database_url: &str, schema: Option<&str>, @@ -915,6 +1226,31 @@ impl PostgresRegistry { row.as_ref().map(map_yaml_import_job).transpose() } + + async fn list_agent_bindings( + &self, + agent_id: &AgentId, + version: u32, + ) -> Result, RegistryError> { + let rows = sqlx::query( + "select + operation_id, + operation_version, + tool_name, + tool_title, + tool_description_override, + enabled + from agent_operation_bindings + where agent_id = $1 and agent_version = $2 + order by tool_name asc", + ) + .bind(agent_id.as_str()) + .bind(to_db_version(version)) + .fetch_all(&self.pool) + .await?; + + rows.iter().map(map_agent_binding).collect() + } } async fn insert_version_row( @@ -968,6 +1304,77 @@ async fn insert_version_row( Ok(()) } +async fn insert_agent_version_row( + tx: &mut Transaction<'_, Postgres>, + version: &AgentVersion, +) -> Result<(), RegistryError> { + sqlx::query( + "insert into agent_versions ( + agent_id, + version, + status, + instructions_json, + tool_selection_policy_json, + created_at + ) values ( + $1, $2, $3, $4, $5, $6::timestamptz + )", + ) + .bind(version.agent_id.as_str()) + .bind(to_db_version(version.version)) + .bind(serialize_enum_text(&version.status, "status")?) + .bind(Json(version.instructions.clone())) + .bind(Json(version.tool_selection_policy.clone())) + .bind(&version.created_at) + .execute(&mut **tx) + .await?; + + Ok(()) +} + +async fn replace_agent_bindings_rows( + tx: &mut Transaction<'_, Postgres>, + agent_id: &AgentId, + version: u32, + bindings: &[AgentOperationBinding], +) -> Result<(), RegistryError> { + sqlx::query( + "delete from agent_operation_bindings + where agent_id = $1 and agent_version = $2", + ) + .bind(agent_id.as_str()) + .bind(to_db_version(version)) + .execute(&mut **tx) + .await?; + + for binding in bindings { + sqlx::query( + "insert into agent_operation_bindings ( + agent_id, + agent_version, + operation_id, + operation_version, + tool_name, + tool_title, + tool_description_override, + enabled + ) values ($1, $2, $3, $4, $5, $6, $7, $8)", + ) + .bind(agent_id.as_str()) + .bind(to_db_version(version)) + .bind(binding.operation_id.as_str()) + .bind(to_db_version(binding.operation_version)) + .bind(&binding.tool_name) + .bind(&binding.tool_title) + .bind(binding.tool_description_override.as_deref()) + .bind(binding.enabled) + .execute(&mut **tx) + .await?; + } + + Ok(()) +} + fn assert_immutable_fields( summary: &OperationSummary, snapshot: &RegistryOperation, @@ -1010,6 +1417,28 @@ fn map_workspace_record(row: &PgRow) -> Result { }) } +fn map_agent_summary(row: &PgRow) -> Result { + Ok(AgentSummary { + id: AgentId::new(row.try_get::("id")?), + workspace_id: WorkspaceId::new(row.try_get::("workspace_id")?), + slug: row.try_get("slug")?, + display_name: row.try_get("display_name")?, + description: row.try_get("description")?, + status: deserialize_enum_text(&row.try_get::("status")?, "status")?, + current_draft_version: from_db_version( + row.try_get("current_draft_version")?, + "current_draft_version", + )?, + latest_published_version: row + .try_get::, _>("latest_published_version")? + .map(|value| from_db_version(value, "latest_published_version")) + .transpose()?, + created_at: row.try_get("created_at")?, + updated_at: row.try_get("updated_at")?, + published_at: row.try_get("published_at")?, + }) +} + fn map_operation_summary(row: &PgRow) -> Result { Ok(OperationSummary { id: OperationId::new(row.try_get::("id")?), @@ -1102,6 +1531,60 @@ fn map_auth_profile(row: &PgRow) -> Result { }) } +fn map_agent_binding(row: &PgRow) -> Result { + Ok(AgentOperationBinding { + operation_id: OperationId::new(row.try_get::("operation_id")?), + operation_version: from_db_version(row.try_get("operation_version")?, "operation_version")?, + tool_name: row.try_get("tool_name")?, + tool_title: row.try_get("tool_title")?, + tool_description_override: row.try_get("tool_description_override")?, + enabled: row.try_get("enabled")?, + }) +} + +fn map_agent_version_record( + summary: &AgentSummary, + row: &PgRow, + bindings: Vec, +) -> Result { + let version = from_db_version(row.try_get("version")?, "version")?; + let status = deserialize_enum_text(&row.try_get::("status")?, "status")?; + + Ok(AgentVersionRecord { + agent_id: summary.id.clone(), + workspace_id: summary.workspace_id.clone(), + version, + status, + created_at: row.try_get("created_at")?, + bindings, + snapshot: AgentVersion { + agent_id: summary.id.clone(), + version, + status, + instructions: row.try_get::, _>("instructions_json")?.0, + tool_selection_policy: row + .try_get::, _>("tool_selection_policy_json")? + .0, + created_at: row.try_get("created_at")?, + }, + }) +} + +fn map_published_agent_tool(row: &PgRow) -> Result { + let record = map_operation_version_record(row)?; + + Ok(PublishedAgentTool { + workspace_id: WorkspaceId::new(row.try_get::("workspace_id")?), + workspace_slug: row.try_get("workspace_slug")?, + agent_id: AgentId::new(row.try_get::("agent_id")?), + agent_slug: row.try_get("agent_slug")?, + tool_name: row.try_get("tool_name")?, + tool_title: row.try_get("tool_title")?, + tool_description: row.try_get("tool_description")?, + operation: record.snapshot, + }) +} + fn map_sample_metadata(row: &PgRow) -> Result { Ok(OperationSampleMetadata { id: crank_core::SampleId::new(row.try_get::("id")?),