diff --git a/scripts/functions-runner.py b/scripts/functions-runner.py index d8ae40fd..a7d8d842 100644 --- a/scripts/functions-runner.py +++ b/scripts/functions-runner.py @@ -98,6 +98,20 @@ def normalize_project_selector(project: Any) -> tuple[str | None, str | None]: return None, None +def project_selector_fields(project: Any) -> dict[str, str]: + project_id, project_name = normalize_project_selector(project) + if project_id: + return {"project_id": project_id} + if not project_name: + return {} + fields = {"project_name": project_name} + # braintrust.framework2.Project exposes `.project_group_name` on SDKs that support project groups. + project_group_name = getattr(project, "project_group_name", None) + if isinstance(project_group_name, str) and project_group_name.strip(): + fields["project_group_name"] = project_group_name.strip() + return fields + + def normalize_function_type(raw: Any) -> str | None: if isinstance(raw, str): value = raw.strip() @@ -189,8 +203,6 @@ def collect_code_entries(functions_registry: Any) -> list[dict[str, Any]]: if not isinstance(name, str) or not isinstance(slug, str) or not name or not slug: continue - project_id, project_name = normalize_project_selector(getattr(item, "project", None)) - entry: dict[str, Any] = { "kind": "code", "name": name, @@ -232,10 +244,7 @@ def collect_code_entries(functions_registry: Any) -> list[dict[str, Any]]: normalized_tags = [tag for tag in tags if isinstance(tag, str)] if normalized_tags: entry["tags"] = normalized_tags - if project_id: - entry["project_id"] = project_id - if project_name: - entry["project_name"] = project_name + entry.update(project_selector_fields(getattr(item, "project", None))) preview = getattr(item, "preview", None) if isinstance(preview, str): @@ -275,13 +284,13 @@ async def collect_function_event_entries(prompts_registry: Any) -> list[dict[str if isinstance(normalized, dict): if normalized.get("if_exists") is None: normalized.pop("if_exists", None) - project_id, project_name = normalize_project_selector(getattr(item, "project", None)) - event_entry: dict[str, Any] = {"kind": "function_event", "event": normalized} - if project_id: - event_entry["project_id"] = project_id - if project_name: - event_entry["project_name"] = project_name - entries.append(event_entry) + entries.append( + { + "kind": "function_event", + "event": normalized, + **project_selector_fields(getattr(item, "project", None)), + } + ) return entries diff --git a/scripts/functions-runner.ts b/scripts/functions-runner.ts index a8a37ceb..03f6599a 100644 --- a/scripts/functions-runner.ts +++ b/scripts/functions-runner.ts @@ -56,6 +56,7 @@ type CodeEntry = { kind: "code"; project_id?: string; project_name?: string; + project_group_name?: string; name: string; slug: string; description?: string; @@ -72,6 +73,7 @@ type FunctionEventEntry = { kind: "function_event"; project_id?: string; project_name?: string; + project_group_name?: string; event: JsonValue; }; @@ -374,6 +376,7 @@ async function collectFunctionEvents( kind: "function_event", project_id: projectId, project_name: projectName, + project_group_name: selector.project_group_name, event: normalizedEvent, }); } @@ -439,6 +442,7 @@ async function collectLegacyPromptEvent( kind: "function_event", project_id: projectId, project_name: projectName, + project_group_name: selector.project_group_name, event, }; } @@ -526,6 +530,7 @@ function collectCodeEntries(items: CodeRegistryItem[]): CodeEntry[] { typeof selector.project_name === "string" ? selector.project_name : undefined, + project_group_name: selector.project_group_name, name: item.name, slug: item.slug, description: diff --git a/scripts/runner-common.ts b/scripts/runner-common.ts index 8a0dd633..4792a986 100644 --- a/scripts/runner-common.ts +++ b/scripts/runner-common.ts @@ -6,11 +6,13 @@ export type JsonValue = JsonPrimitive | JsonArray | JsonObject; export type ProjectSelector = { project_id?: string; project_name?: string; + project_group_name?: string; }; export type ProjectRef = { id?: string; name?: string; + projectGroupName?: string; }; export function asProjectSelector( @@ -25,7 +27,13 @@ export function asProjectSelector( } if (typeof project.name === "string" && project.name.trim().length > 0) { - return { project_name: project.name }; + return typeof project.projectGroupName === "string" && + project.projectGroupName.trim().length > 0 + ? { + project_name: project.name, + project_group_name: project.projectGroupName, + } + : { project_name: project.name }; } return {}; diff --git a/src/datasets/pipeline.rs b/src/datasets/pipeline.rs index 9a04e014..ce0565a1 100644 --- a/src/datasets/pipeline.rs +++ b/src/datasets/pipeline.rs @@ -1808,7 +1808,7 @@ async fn resolve_target_project( if let Some(project) = get_project_by_name(client, project_name).await? { Ok(project) } else { - create_project(client, project_name) + create_project(client, project_name, None) .await .with_context(|| format!("project '{project_name}' not found, and creating it failed")) } diff --git a/src/functions/push.rs b/src/functions/push.rs index 72c8ee3e..5d253e08 100644 --- a/src/functions/push.rs +++ b/src/functions/push.rs @@ -1,4 +1,4 @@ -use std::collections::{BTreeMap, BTreeSet}; +use std::collections::{btree_map, BTreeMap, BTreeSet}; use std::ffi::OsString; use std::path::{Path, PathBuf}; use std::process::{Command, Output}; @@ -97,6 +97,8 @@ struct CodeEntry { project_id: Option, #[serde(default)] project_name: Option, + #[serde(default)] + project_group_name: Option, name: String, slug: String, #[serde(default)] @@ -123,6 +125,8 @@ struct FunctionEventEntry { project_id: Option, #[serde(default)] project_name: Option, + #[serde(default)] + project_group_name: Option, event: Value, } @@ -201,6 +205,8 @@ struct ProjectPreflight { default_project_name: Option, requires_default_project: bool, named_projects: BTreeSet, + /// Project group to create each named project in, keyed by project name. + project_group_names: BTreeMap, direct_project_ids: BTreeSet, } @@ -539,6 +545,7 @@ pub async fn run(base: BaseArgs, args: PushArgs) -> Result<()> { let mut project_name_cache = match resolve_named_projects( &auth_ctx, &preflight.named_projects, + &preflight.project_group_names, args.create_missing_projects, ) .await @@ -2550,19 +2557,48 @@ fn collect_project_preflight( let default_project_name = resolve_default_project_name(base, Some(resolved_org))?; let mut requires_default_project = false; let mut named_projects = BTreeSet::new(); + let mut project_group_names = BTreeMap::new(); let mut direct_project_ids = BTreeSet::new(); for file in &manifest.files { for entry in &file.entries { - let selector = match entry { - ManifestEntry::Code(code) => project_selector_for_code(code)?, + let (selector, project_group_name) = match entry { + ManifestEntry::Code(code) => ( + project_selector_for_code(code)?, + code.project_group_name.as_deref(), + ), ManifestEntry::FunctionEvent(event) => { let mut placeholders = BTreeSet::new(); collect_project_name_placeholders_checked(&event.event, &mut placeholders)?; named_projects.extend(placeholders); - project_selector_for_event(event)? + ( + project_selector_for_event(event)?, + event.project_group_name.as_deref(), + ) } }; + if let (ProjectSelector::Name(project_name), Some(project_group_name)) = ( + &selector, + project_group_name + .map(str::trim) + .filter(|value| !value.is_empty()), + ) { + match project_group_names.entry(project_name.clone()) { + btree_map::Entry::Vacant(slot) => { + slot.insert(project_group_name.to_string()); + } + btree_map::Entry::Occupied(existing) + if existing.get() != project_group_name => + { + bail!( + "project '{project_name}' is assigned to conflicting project groups '{}' and '{project_group_name}'", + existing.get() + ); + } + btree_map::Entry::Occupied(_) => {} + } + } + add_selector_requirement( file, entry_slug(entry)?, @@ -2579,6 +2615,7 @@ fn collect_project_preflight( default_project_name, requires_default_project, named_projects, + project_group_names, direct_project_ids, }) } @@ -2721,6 +2758,7 @@ fn resolve_default_project_id( async fn resolve_named_projects( auth_ctx: &super::AuthContext, named_projects: &BTreeSet, + project_group_names: &BTreeMap, create_missing_projects: bool, ) -> Result> { let mut project_name_cache = BTreeMap::new(); @@ -2737,17 +2775,18 @@ async fn resolve_named_projects( continue; } - match create_project(&auth_ctx.client, project_name).await { + let project_group_name = project_group_names.get(project_name).map(String::as_str); + match create_project(&auth_ctx.client, project_name, project_group_name).await { Ok(project) => { project_name_cache.insert(project_name.clone(), project.id); } - Err(_) => { + Err(err) => { // Another writer may have created the project concurrently. if let Some(project) = get_project_by_name(&auth_ctx.client, project_name).await? { project_name_cache.insert(project_name.clone(), project.id); } else { bail!( - "failed to create project '{project_name}' in org '{}'", + "failed to create project '{project_name}' in org '{}': {err:#}", current_org_label(auth_ctx) ); } @@ -2957,7 +2996,8 @@ async fn resolve_project_name( let project = if let Some(project) = get_project_by_name(client, project_name).await? { project } else if create_missing_projects { - match create_project(client, project_name).await { + // Named projects, including their project groups, are created during preflight. + match create_project(client, project_name, None).await { Ok(project) => project, Err(_) => get_project_by_name(client, project_name) .await? @@ -3249,6 +3289,7 @@ mod tests { ManifestEntry::Code(CodeEntry { project_id: None, project_name: None, + project_group_name: None, name: "Code Tool".to_string(), slug: "code-tool".to_string(), description: None, @@ -3263,6 +3304,7 @@ mod tests { ManifestEntry::FunctionEvent(FunctionEventEntry { project_id: None, project_name: None, + project_group_name: None, event: serde_json::json!({ "name": " Prompt Function ", "slug": " prompt-function " @@ -3368,6 +3410,7 @@ mod tests { entries: vec![ManifestEntry::Code(CodeEntry { project_id: None, project_name: None, + project_group_name: None, name: "A".to_string(), slug: "same".to_string(), description: None, @@ -3392,6 +3435,69 @@ mod tests { ); } + fn group_code_entry(slug: &str, project_group_name: &str) -> ManifestEntry { + ManifestEntry::Code(CodeEntry { + project_id: None, + project_name: Some("test-project".to_string()), + project_group_name: Some(project_group_name.to_string()), + name: slug.to_string(), + slug: slug.to_string(), + description: None, + function_type: Some("tool".to_string()), + if_exists: None, + metadata: None, + tags: None, + function_schema: None, + location: None, + preview: None, + }) + } + + fn group_manifest(entries: Vec) -> RunnerManifest { + RunnerManifest { + runtime_context: RuntimeContext { + runtime: "node".to_string(), + version: "20.0.0".to_string(), + }, + files: vec![ManifestFile { + source_file: "a.ts".to_string(), + entries, + python_bundle: None, + }], + baseline_dep_versions: vec![], + } + } + + #[test] + fn collect_project_preflight_records_project_groups() { + let manifest = group_manifest(vec![ + group_code_entry("a", "test-group"), + group_code_entry("b", "test-group"), + ]); + + let preflight = + collect_project_preflight(&test_base_args(), &manifest, "test-org").expect("preflight"); + assert_eq!( + preflight.project_group_names.get("test-project"), + Some(&"test-group".to_string()) + ); + } + + #[test] + fn collect_project_preflight_rejects_conflicting_project_groups() { + let manifest = group_manifest(vec![ + group_code_entry("a", "test-group-a"), + group_code_entry("b", "test-group-b"), + ]); + + let err = collect_project_preflight(&test_base_args(), &manifest, "test-org") + .expect_err("conflicting groups must fail"); + assert!( + err.to_string().contains("conflicting project groups"), + "unexpected error: {err}" + ); + } + #[test] fn explicit_org_validation_rejects_unknown_org() { let mut base = test_base_args(); @@ -3652,6 +3758,7 @@ mod tests { entries: vec![ManifestEntry::Code(CodeEntry { project_id: None, project_name: None, + project_group_name: None, name: "Tool".to_string(), slug: "tool".to_string(), description: None, @@ -3697,6 +3804,7 @@ mod tests { entries: vec![ManifestEntry::Code(CodeEntry { project_id: None, project_name: None, + project_group_name: None, name: "Tool".to_string(), slug: "tool".to_string(), description: None, @@ -3743,6 +3851,7 @@ mod tests { entries: vec![ManifestEntry::Code(CodeEntry { project_id: None, project_name: None, + project_group_name: None, name: "Tool".to_string(), slug: "tool".to_string(), description: None, @@ -3792,6 +3901,7 @@ mod tests { entries: vec![ManifestEntry::Code(CodeEntry { project_id: None, project_name: None, + project_group_name: None, name: "Tool".to_string(), slug: "tool".to_string(), description: None, diff --git a/src/projects/api.rs b/src/projects/api.rs index fc92703d..9ac555ab 100644 --- a/src/projects/api.rs +++ b/src/projects/api.rs @@ -24,8 +24,15 @@ pub async fn list_projects(client: &ApiClient) -> Result> { Ok(list.objects) } -pub async fn create_project(client: &ApiClient, name: &str) -> Result { - let body = serde_json::json!({ "name": name, "org_name": client.org_name() }); +pub async fn create_project( + client: &ApiClient, + name: &str, + project_group_name: Option<&str>, +) -> Result { + let mut body = serde_json::json!({ "name": name, "org_name": client.org_name() }); + if let Some(project_group_name) = project_group_name { + body["project_group_name"] = project_group_name.into(); + } client.post("/v1/project", &body).await } diff --git a/src/projects/create.rs b/src/projects/create.rs index 69946562..7fdd4c81 100644 --- a/src/projects/create.rs +++ b/src/projects/create.rs @@ -32,7 +32,7 @@ pub(crate) async fn create_project_checked( match with_spinner_visible( "Creating project...", - api::create_project(client, name), + api::create_project(client, name, None), Duration::from_millis(300), ) .await diff --git a/src/switch.rs b/src/switch.rs index af2adad8..8a04a74e 100644 --- a/src/switch.rs +++ b/src/switch.rs @@ -282,7 +282,11 @@ async fn validate_or_create_project(client: &ApiClient, name: &str) -> Result = + serde_json::from_str(stdout.trim()).expect("parse entries JSON from project group script"); + let project_group_names: Vec> = entries + .iter() + .map(|entry| entry.get("project_group_name").and_then(Value::as_str)) + .collect(); + assert_eq!( + project_group_names, + vec![Some("test-group"), None, Some("test-group")] + ); + assert!(entries + .iter() + .all(|entry| entry.get("project_name").and_then(Value::as_str) == Some("test-project"))); +} + #[test] fn functions_js_runner_emits_valid_manifest() { if !command_exists("node") {