diff --git a/src/assets/templates/a2a-python-strands/pyproject.toml b/src/assets/templates/a2a-python-strands/pyproject.toml index a7b9ac465..85aa6b900 100644 --- a/src/assets/templates/a2a-python-strands/pyproject.toml +++ b/src/assets/templates/a2a-python-strands/pyproject.toml @@ -9,15 +9,16 @@ description = "AgentCore A2A Agent using Strands SDK" readme = "README.md" requires-python = ">=3.10" dependencies = [ - {{#if (eq modelProvider "Anthropic")}}"anthropic ~= 0.30.0", - {{/if}}"a2a-sdk[all] >= 0.3.0, < 0.4.0", + "a2a-sdk[all] ~= 0.3.0", "aws-opentelemetry-distro ~= 0.17.0", "bedrock-agentcore[a2a] ~= 1.9.1", "botocore[crt] ~= 1.43.0", - {{#if (eq modelProvider "Gemini")}}"google-genai ~= 1.0.0", - {{/if}}{{#if (eq modelProvider "OpenAI")}}"openai ~= 1.0.0", - {{/if}}{{#if (eq modelProvider "LiteLLM")}}"litellm ~= 1.0.0", - {{/if}}"strands-agents ~= 1.15.0", + {{#if (eq modelProvider "Anthropic")}}"strands-agents[anthropic] ~= 1.15.0", + {{else}}{{#if (eq modelProvider "OpenAI")}}"strands-agents[openai] ~= 1.15.0", + {{else}}{{#if (eq modelProvider "Gemini")}}"strands-agents[gemini] ~= 1.15.0", + {{else}}{{#if (eq modelProvider "LiteLLM")}}"strands-agents[litellm] ~= 1.15.0", + {{else}}"strands-agents ~= 1.15.0", + {{/if}}{{/if}}{{/if}}{{/if}} ] [tool.hatch.build.targets.wheel] diff --git a/src/assets/templates/agent-python-strands/README.md b/src/assets/templates/agent-python-strands/README.md index 5714aafbf..e3edfd5ae 100644 --- a/src/assets/templates/agent-python-strands/README.md +++ b/src/assets/templates/agent-python-strands/README.md @@ -23,7 +23,7 @@ invoking the agent. | Variable | Required | Description | | --- | --- | --- | -{{#if hasIdentity}}| `{{identityProviders.[0].envVarName}}` | Yes | {{modelProvider}} API key (local) or Identity provider name (deployed) | +{{#if identityProviders.[0]}}| `{{identityProviders.[0].envVarName}}` | Yes | {{modelProvider}} API key (local) or Identity provider name (deployed) | {{/if}}| `LOCAL_DEV` | No | Set to `1` to use `.env.local` instead of AgentCore Identity | # Developing locally diff --git a/src/assets/templates/agent-python-strands/pyproject.toml b/src/assets/templates/agent-python-strands/pyproject.toml index 0d3a70143..88eadd517 100644 --- a/src/assets/templates/agent-python-strands/pyproject.toml +++ b/src/assets/templates/agent-python-strands/pyproject.toml @@ -9,17 +9,18 @@ description = "AgentCore Runtime Application using Strands SDK" readme = "README.md" requires-python = ">=3.10" dependencies = [ - {{#if (eq modelProvider "Anthropic")}}"anthropic ~= 0.30.0", - {{/if}}"aws-opentelemetry-distro ~= 0.18.0", + "aws-opentelemetry-distro ~= 0.18.0", "bedrock-agentcore ~= 1.9.1", "botocore[crt] ~= 1.43.0", - {{#if (eq modelProvider "Gemini")}}"google-genai ~= 1.0.0", - {{/if}}"mcp ~= 1.24.0", - {{#if (eq modelProvider "OpenAI")}}"openai ~= 1.0.0", - {{/if}}{{#if (eq modelProvider "LiteLLM")}}"litellm ~= 1.0.0", - {{/if}}{{#if bedrockMantle}}"openai ~= 1.0.0", + "mcp ~= 1.24.0", + {{#if bedrockMantle}}"openai ~= 1.0.0", "aws-bedrock-token-generator ~= 1.0.0", - {{/if}}"strands-agents ~= 1.15.0", + {{/if}}{{#if (eq modelProvider "Anthropic")}}"strands-agents[anthropic] ~= 1.15.0", + {{else}}{{#if (eq modelProvider "OpenAI")}}"strands-agents[openai] ~= 1.15.0", + {{else}}{{#if (eq modelProvider "Gemini")}}"strands-agents[gemini] ~= 1.15.0", + {{else}}{{#if (eq modelProvider "LiteLLM")}}"strands-agents[litellm] ~= 1.15.0", + {{else}}"strands-agents ~= 1.15.0", + {{/if}}{{/if}}{{/if}}{{/if}} {{#if (or hasBrowser hasCodeInterpreter)}}"strands-agents-tools ~= 0.1.0", {{/if}}{{#if hasBrowser}}"nest-asyncio ~= 1.5.0", "playwright ~= 1.42.0", diff --git a/src/assets/templates/agent-typescript-strands/README.md b/src/assets/templates/agent-typescript-strands/README.md index 08f22a6fd..4c64ecd57 100644 --- a/src/assets/templates/agent-typescript-strands/README.md +++ b/src/assets/templates/agent-typescript-strands/README.md @@ -22,7 +22,7 @@ validation when extending the request shape, and pass only prompt text to the ag | Variable | Required | Description | | --- | --- | --- | -{{#if hasIdentity}}| `{{identityProviders.[0].envVarName}}` | Yes | {{modelProvider}} API key (local) or Identity provider name (deployed) | +{{#if identityProviders.[0]}}| `{{identityProviders.[0].envVarName}}` | Yes | {{modelProvider}} API key (local) or Identity provider name (deployed) | {{/if}}| `LOCAL_DEV` | No | Set to `1` to use `.env.local` instead of AgentCore Identity | # Developing locally diff --git a/src/assets/templates/agent-typescript-strands/package.json b/src/assets/templates/agent-typescript-strands/package.json.template similarity index 68% rename from src/assets/templates/agent-typescript-strands/package.json rename to src/assets/templates/agent-typescript-strands/package.json.template index 77eb3bc8e..804fec0f1 100644 --- a/src/assets/templates/agent-typescript-strands/package.json +++ b/src/assets/templates/agent-typescript-strands/package.json.template @@ -10,11 +10,14 @@ "dev": "tsx watch main.ts" }, "dependencies": { - "@modelcontextprotocol/sdk": "~1.25.2", + {{#if (eq modelProvider "Anthropic")}}"@anthropic-ai/sdk": "~0.92.0", + {{/if}}{{#if (eq modelProvider "Gemini")}}"@google/genai": "~1.40.0", + {{/if}}"@modelcontextprotocol/sdk": "~1.25.2", "@opentelemetry/api": "~1.9.0", "@strands-agents/sdk": "~1.5.0", "bedrock-agentcore": "~0.3.0", - "tsx": "~4.19.0", + {{#if (eq modelProvider "OpenAI")}}"openai": "~6.7.0", + {{/if}}"tsx": "~4.19.0", "zod": "~4.4.3" }, "devDependencies": { diff --git a/src/core/project/manager.tsx b/src/core/project/manager.tsx index c502cbdae..d001a7e89 100644 --- a/src/core/project/manager.tsx +++ b/src/core/project/manager.tsx @@ -169,13 +169,18 @@ export class FsProjectManager implements ProjectManager { const destination = join(process.cwd(), input.name); yield { type: "step", message: "Creating project tree" }; - const projectTree = await createProjectTree( + const { tree: projectTree, envEntries } = await createProjectTree( { templateRenderer: this.templateRenderer, assetSource: this.assetSource }, { projectName: input.name }, { runtime: scaffoldRuntimeInput, importBedrockAgent: input.importBedrockAgent }, ); await projectTree.write(destination); + if (envEntries.length > 0) { + yield { type: "step", message: "Writing model provider API key to agentcore/.env.local" }; + await new EnvLocalFile(destination).insertIfNew(envEntries); + } + // A harness project scaffolds through the same addResource flow that // `project add harness` uses, so a create-time harness and an added one can // never drift apart. @@ -306,10 +311,24 @@ export class FsProjectManager implements ProjectManager { const outputPath = join(project.rootPath, "app", input.resourceConfig.name); scaffoldedPaths.push(outputPath); - const spec = await this.scaffoldRuntimeResources(outputPath, input.resourceConfig); + const { spec, envEntries } = await this.scaffoldRuntimeResources( + outputPath, + input.resourceConfig, + ); if (spec.runtimes) projectSpec.runtimes.push(...spec.runtimes); if (spec.memories) projectSpec.memories.push(...spec.memories); if (spec.credentials) projectSpec.credentials.push(...spec.credentials); + if (envEntries.length > 0) { + envFile = new EnvLocalFile(project.rootPath); + yield { type: "step", message: `Updating secrets file at '${envFile.path}'` }; + const { skipped } = await envFile.insertIfNew(envEntries); + for (const key of skipped) { + yield { + type: "step", + message: `'${key}' already exists in ${ENV_LOCAL_RELATIVE_PATH}; left unchanged`, + }; + } + } yield* this.installRuntimeDependencies(outputPath); break; @@ -834,7 +853,7 @@ export class FsProjectManager implements ProjectManager { const result = await resolver.resolve(input); await result.tree.write(dirname(outputPath)); - return result.spec; + return { spec: result.spec, envEntries: result.envEntries ?? [] }; } public async *build(project: Project): AsyncGenerator { diff --git a/src/core/project/templates/project.ts b/src/core/project/templates/project.ts index c2eeab8e4..980e3b865 100644 --- a/src/core/project/templates/project.ts +++ b/src/core/project/templates/project.ts @@ -7,7 +7,9 @@ import type { } from "../../../handlers/project/add/runtime/types"; import { InputValidationError } from "../../../errors/errors"; import { getRuntimeTemplateResolver } from "./runtime"; -import type { SpecEntries, Template, TemplateRenderer } from "./types"; +import { mergeSpecEntries } from "./spec"; +import type { Template, TemplateRenderer } from "./types"; +import type { EnvLocalEntry } from "../../../handlers/project/types"; type CreateProjectConfig = { assetSource: AssetSource; @@ -18,7 +20,7 @@ export async function createProjectTree( config: CreateProjectConfig, input: { projectName: string }, options?: { runtime?: ScaffoldRuntimeInput; importBedrockAgent?: ImportBedrockAgentInput }, -): Promise { +): Promise<{ tree: FsTreeNode; envEntries: EnvLocalEntry[] }> { const templates: Template[] = []; if (options?.runtime) { const runtimeConfig: RuntimeResourceConfig = { @@ -34,7 +36,9 @@ export async function createProjectTree( templates.push(await resolver.resolve(runtimeConfig)); } - return FsTreeNode.createDirectory(".", [ + const envEntries = templates.flatMap((template) => template.envEntries ?? []); + + const tree = FsTreeNode.createDirectory(".", [ FsTreeNode.createFile(".gitignore", () => config.assetSource.read("templates/shared/gitignore.template"), ), @@ -58,20 +62,8 @@ export async function createProjectTree( templates.map((t) => t.tree), ), ]); + + return { tree, envEntries }; } const json = (value: unknown): string => `${JSON.stringify(value, null, 2)}\n`; - -function mergeSpecEntries(entries: SpecEntries[]): SpecEntries { - const runtimes = entries.flatMap(({ runtimes }) => runtimes ?? []); - const credentials = entries.flatMap(({ credentials }) => credentials ?? []); - const memories = entries.flatMap(({ memories }) => memories ?? []); - const harnesses = entries.flatMap(({ harnesses }) => harnesses ?? []); - - return { - ...(runtimes.length > 0 && { runtimes }), - ...(credentials.length > 0 && { credentials }), - ...(memories.length > 0 && { memories }), - ...(harnesses.length > 0 && { harnesses }), - }; -} diff --git a/src/core/project/templates/runtime.ts b/src/core/project/templates/runtime.ts index 98c733ff8..01c5b2ba7 100644 --- a/src/core/project/templates/runtime.ts +++ b/src/core/project/templates/runtime.ts @@ -2,11 +2,40 @@ import { FsTreeNode } from "./fsTree"; import type { AssetSource } from "../source"; import type { RuntimeResourceConfig } from "../../../handlers/project/add/runtime/types"; import type { ProjectRuntime } from "../../../projectSchemas/runtime"; -import type { TemplateRenderer, TemplateResolver } from "./types"; -import type { ScaffoldRuntimeInput } from "../../../handlers/project/types"; +import { mergeSpecEntries } from "./spec"; +import type { SpecEntries, TemplateRenderer, TemplateResolver } from "./types"; +import type { EnvLocalEntry, ScaffoldRuntimeInput } from "../../../handlers/project/types"; +import { credentialEnvVarName } from "../../../projectSchemas/credential"; import { InputValidationError } from "../../../errors"; import { toPythonPackageName } from "../fsUtils"; +/** A model provider's render context, spec entries, and .env.local secrets for a scaffolded runtime. */ +type ModelProviderTemplateConfig = { + templateRenderContext: { identityProviders: { name: string; envVarName: string }[] }; + spec: SpecEntries; + envEntries: EnvLocalEntry[]; +}; + +function resolveModelProviderScaffold(input: RuntimeResourceConfig): ModelProviderTemplateConfig { + const { modelProvider, apiKey } = input.scaffoldRuntimeInput; + if (apiKey === undefined) { + return { templateRenderContext: { identityProviders: [] }, spec: {}, envEntries: [] }; + } + const credentialName = `${input.name}${modelProvider}ApiKey`; + const envVarName = credentialEnvVarName(credentialName); + return { + templateRenderContext: { identityProviders: [{ name: credentialName, envVarName }] }, + spec: { credentials: [{ authorizerType: "ApiKeyCredentialProvider", name: credentialName }] }, + envEntries: [ + { + key: envVarName, + value: apiKey, + comment: `API key for the ${modelProvider} model provider (runtime ${input.name})`, + }, + ], + }; +} + function buildRuntimeSpec(input: RuntimeResourceConfig): ProjectRuntime { const { scaffoldRuntimeInput, name, ...infra } = input; return { @@ -89,6 +118,11 @@ const importBedrockAgentResolver = () => async (input: RuntimeResourceConfig) => const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: TemplateRenderer) => ({ [buildResolverKey("none", "Python", "HTTP")]: async (input: RuntimeResourceConfig) => { + const { modelProvider } = input.scaffoldRuntimeInput; + if (modelProvider !== undefined && modelProvider !== "Bedrock") + throw new InputValidationError( + "the agent-python template only supports the Bedrock model provider", + ); if (input.scaffoldRuntimeInput.memory !== undefined) throw new InputValidationError(`memory is not supported with the agent-python template`); const isContainer = input.scaffoldRuntimeInput.build === "Container"; @@ -121,18 +155,18 @@ const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: Templa : [], ); const memory = input.scaffoldRuntimeInput.memory; + const modelScaffold = resolveModelProviderScaffold(input); const context = { name: toPythonPackageName(input.name), - modelProvider: input.scaffoldRuntimeInput.modelProvider, + modelProvider: input.scaffoldRuntimeInput.modelProvider ?? "Bedrock", hasMemory: memory !== undefined, // the CDK injects this env var corresponding to the actual ID once its resolved on deployment. memoryEnvVarName: memory ? `MEMORY_${memory.name.toUpperCase()}_ID` : undefined, memoryStrategies: memory?.strategies.map(({ type }) => type) ?? [], - hasIdentity: false, + ...modelScaffold.templateRenderContext, hasGateway: false, hasPayment: false, isVpc: input.networkMode === "VPC", - identityProviders: [], gatewayProviders: [], gatewayAuthTypes: [], sessionStorageMountPath, @@ -160,15 +194,23 @@ const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: Templa ); return { tree, - spec: { - runtimes: [{ ...buildRuntimeSpec(input), protocol: "HTTP" as const }], - ...(memory && { memories: [memory] }), - }, + spec: mergeSpecEntries([ + { + runtimes: [{ ...buildRuntimeSpec(input), protocol: "HTTP" as const }], + ...(memory && { memories: [memory] }), + }, + modelScaffold.spec, + ]), + ...(modelScaffold.envEntries.length > 0 && { envEntries: modelScaffold.envEntries }), }; }, [buildResolverKey("strands", "TypeScript", "HTTP")]: async (input: RuntimeResourceConfig) => { if (input.protocol !== undefined && input.protocol !== "HTTP") throw new InputValidationError("the agent-typescript-strands template only supports HTTP"); + if (input.scaffoldRuntimeInput.modelProvider === "LiteLLM") + throw new InputValidationError( + "the agent-typescript-strands template does not support the LiteLLM model provider", + ); const memory = input.scaffoldRuntimeInput.memory; // The TypeScript strands SDK's createAgentCoreMemoryStores requires at least one @@ -179,15 +221,15 @@ const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: Templa "the agent-typescript-strands template does not support short-term-only memory; add long-term strategies or use --memory none", ); + const modelScaffold = resolveModelProviderScaffold(input); const context = { name: toNpmPackageName(input.name), - modelProvider: input.scaffoldRuntimeInput.modelProvider, + modelProvider: input.scaffoldRuntimeInput.modelProvider ?? "Bedrock", hasMemory: memory !== undefined, // the CDK injects this env var corresponding to the actual ID once its resolved on deployment. memoryEnvVarName: memory ? `MEMORY_${memory.name.toUpperCase()}_ID` : undefined, memoryStrategies: memory?.strategies.map(({ type }) => type) ?? [], - hasIdentity: false, - identityProviders: [], + ...modelScaffold.templateRenderContext, }; const isContainer = input.scaffoldRuntimeInput.build === "Container"; const tree = await FsTreeNode.fromAssetSource( @@ -205,13 +247,19 @@ const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: Templa ); return { tree, - spec: { - runtimes: [{ ...buildRuntimeSpec(input), protocol: "HTTP" as const }], - ...(memory && { memories: [memory] }), - }, + spec: mergeSpecEntries([ + { + runtimes: [{ ...buildRuntimeSpec(input), protocol: "HTTP" as const }], + ...(memory && { memories: [memory] }), + }, + modelScaffold.spec, + ]), + ...(modelScaffold.envEntries.length > 0 && { envEntries: modelScaffold.envEntries }), }; }, [buildResolverKey("none", "Python", "MCP")]: async (input: RuntimeResourceConfig) => { + if (input.scaffoldRuntimeInput.modelProvider !== undefined) + throw new InputValidationError("an MCP runtime does not use a model provider"); if (input.scaffoldRuntimeInput.memory !== undefined) throw new InputValidationError("memory is not supported with an MCP runtime"); const filesystemConfigurations = input.filesystemConfigurations ?? []; @@ -274,13 +322,15 @@ const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: Templa : [], ); const memory = input.scaffoldRuntimeInput.memory; + const modelScaffold = resolveModelProviderScaffold(input); const context = { name: toPythonPackageName(input.name), - modelProvider: input.scaffoldRuntimeInput.modelProvider, + modelProvider: input.scaffoldRuntimeInput.modelProvider ?? "Bedrock", hasMemory: memory !== undefined, // the CDK injects this env var corresponding to the actual ID once its resolved on deployment. memoryEnvVarName: memory ? `MEMORY_${memory.name.toUpperCase()}_ID` : undefined, memoryStrategies: memory?.strategies.map(({ type }) => type) ?? [], + ...modelScaffold.templateRenderContext, sessionStorageMountPath, efsMounts, s3Mounts, @@ -307,10 +357,14 @@ const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: Templa ); return { tree, - spec: { - runtimes: [{ ...buildRuntimeSpec(input), protocol: "A2A" as const }], - ...(memory && { memories: [memory] }), - }, + spec: mergeSpecEntries([ + { + runtimes: [{ ...buildRuntimeSpec(input), protocol: "A2A" as const }], + ...(memory && { memories: [memory] }), + }, + modelScaffold.spec, + ]), + ...(modelScaffold.envEntries.length > 0 && { envEntries: modelScaffold.envEntries }), }; }, }); diff --git a/src/core/project/templates/spec.ts b/src/core/project/templates/spec.ts new file mode 100644 index 000000000..a412ca74f --- /dev/null +++ b/src/core/project/templates/spec.ts @@ -0,0 +1,16 @@ +import type { SpecEntries } from "./types"; + +/** Combines several {@link SpecEntries} into one, concatenating each resource collection. */ +export function mergeSpecEntries(entries: SpecEntries[]): SpecEntries { + const runtimes = entries.flatMap(({ runtimes }) => runtimes ?? []); + const credentials = entries.flatMap(({ credentials }) => credentials ?? []); + const memories = entries.flatMap(({ memories }) => memories ?? []); + const harnesses = entries.flatMap(({ harnesses }) => harnesses ?? []); + + return { + ...(runtimes.length > 0 && { runtimes }), + ...(credentials.length > 0 && { credentials }), + ...(memories.length > 0 && { memories }), + ...(harnesses.length > 0 && { harnesses }), + }; +} diff --git a/src/core/project/templates/types.ts b/src/core/project/templates/types.ts index 798e4ba4b..d7e386301 100644 --- a/src/core/project/templates/types.ts +++ b/src/core/project/templates/types.ts @@ -4,6 +4,7 @@ import type { MemorySchema } from "../../../projectSchemas/memory"; import type { CredentialSchema } from "../../../projectSchemas/credential"; import type { HarnessRegistryEntry } from "../../../projectSchemas/harness"; import type { Evaluator } from "../../../projectSchemas/evaluator"; +import type { EnvLocalEntry } from "../../../handlers/project/types"; import type z from "zod"; /** AgentCore Project Spec Entries that rendered as part of a {@link Template} **/ @@ -19,6 +20,8 @@ export type SpecEntries = { export type Template = { tree: FsTreeNode; spec: SpecEntries; + /** Secret material for agentcore/.env.local (e.g. a model provider API key). */ + envEntries?: EnvLocalEntry[]; }; /** A standard interface for resolving templates from a given input of paramters **/ diff --git a/src/handlers/project/add/runtime/index.test.ts b/src/handlers/project/add/runtime/index.test.ts index 17869ac44..59118c54a 100644 --- a/src/handlers/project/add/runtime/index.test.ts +++ b/src/handlers/project/add/runtime/index.test.ts @@ -11,6 +11,7 @@ import { } from "../../../../testing"; import { InputValidationError } from "../../../../errors"; import type { BedrockAgentImportPlan } from "../../../../core/project/bedrockAgentImport"; +import { credentialEnvVarName } from "../../../../projectSchemas/credential"; const originalCwd = process.cwd(); const tempDirectories: string[] = []; @@ -300,8 +301,6 @@ describe("project add runtime", () => { "none", "--protocol", "MCP", - "--model-provider", - "Bedrock", "--memory", "none", ], @@ -698,12 +697,27 @@ describe("project add runtime", () => { "none", "--protocol", "MCP", - "--model-provider", - "Bedrock", "--memory", "shortTerm", ], ], + [ + "MCP runtime rejects a model provider", + [ + "--name", + "my_agent", + "--build", + "CodeZip", + "--language", + "Python", + "--framework", + "none", + "--protocol", + "MCP", + "--model-provider", + "Bedrock", + ], + ], [ "--memory shortTerm is not supported with --framework none", [ @@ -781,6 +795,80 @@ describe("project add runtime", () => { ]), ).rejects.toThrow(/API keys are not compatible with Bedrock model providers/); }); + + test.each<[string, string, string]>([ + ["anthropic", "Anthropic", "Python"], + ["OpenAI", "OpenAI", "Python"], + ["gemini", "Gemini", "Python"], + ["Anthropic", "Anthropic", "TypeScript"], + ])( + "scaffolds a strands runtime for --model-provider %s (%s) with an API-key credential", + async (flagValue, provider, language) => { + const projectRoot = await inProject(); + const apiKeyPath = join(projectRoot, "api-key.txt"); + await Bun.write(apiKeyPath, "test-api-key"); + + await run([ + "add", + "runtime", + "--name", + "my_agent", + "--build", + "CodeZip", + "--language", + language, + "--framework", + "strands", + "--model-provider", + flagValue, + "--api-key", + `file://${apiKeyPath}`, + ]); + + const credentialName = `my_agent${provider}ApiKey`; + const envVarName = credentialEnvVarName(credentialName); + + const spec = await Bun.file(join(projectRoot, "agentcore", "agentcore.json")).json(); + expect(spec.credentials).toContainEqual({ + authorizerType: "ApiKeyCredentialProvider", + name: credentialName, + }); + + const envLocal = await Bun.file(join(projectRoot, "agentcore", ".env.local")).text(); + expect(envLocal).toContain(`${envVarName}='test-api-key'`); + }, + ); + + test.each<[string, string, string, boolean]>([ + ["Anthropic without an API key", "Anthropic", "strands", false], + ["OpenAI without an API key", "OpenAI", "strands", false], + ["Gemini without an API key", "Gemini", "strands", false], + ["a non-Bedrock provider on a provider-less template", "Anthropic", "none", true], + ["LiteLLM on the TypeScript template", "LiteLLM", "strands", false], + ])("rejects %s", async (_label, provider, framework, includeApiKey) => { + const projectRoot = await inProject(); + const apiKeyPath = join(projectRoot, "api-key.txt"); + await Bun.write(apiKeyPath, "test-api-key"); + + const language = provider === "LiteLLM" ? "TypeScript" : "Python"; + const flags = [ + "add", + "runtime", + "--name", + "my_agent", + "--build", + "CodeZip", + "--language", + language, + "--framework", + framework, + "--model-provider", + provider, + ]; + if (includeApiKey) flags.push("--api-key", `file://${apiKeyPath}`); + + await expect(run(flags)).rejects.toBeInstanceOf(InputValidationError); + }); }); describe("project add runtime --type import", () => { diff --git a/src/handlers/project/add/runtime/index.ts b/src/handlers/project/add/runtime/index.ts index 4f629e7c5..03c6fd2e1 100644 --- a/src/handlers/project/add/runtime/index.ts +++ b/src/handlers/project/add/runtime/index.ts @@ -14,7 +14,11 @@ import { RUNTIME_TEMPLATE_SHORTCUT_NAMES, resolveRuntimeTemplateShortcut, } from "../../shortcuts"; -import { ScaffoldRuntimeInputSchema, type ScaffoldRuntimeInput } from "../../types"; +import { + ModelProviderSchema, + ScaffoldRuntimeInputSchema, + type ScaffoldRuntimeInput, +} from "../../types"; import { RuntimeResourceConfigSchema, type ImportBedrockAgentInput } from "./types"; import { importScaffoldRuntimeInput, @@ -63,8 +67,8 @@ export const createAddRuntimeHandler = (config: AddProjectResourceConfig) => ), flag( "model-provider", - "model provider for the scaffolded runtime code", - z.enum(["Bedrock"]).optional(), + "model provider for the scaffolded runtime code (Bedrock, Anthropic, OpenAI, or Gemini)", + ModelProviderSchema.optional(), ), flag( "api-key", diff --git a/src/handlers/project/create/index.ts b/src/handlers/project/create/index.ts index 4277f34a1..083e2de97 100644 --- a/src/handlers/project/create/index.ts +++ b/src/handlers/project/create/index.ts @@ -11,6 +11,7 @@ import { import { ScaffoldRuntimeInputSchema, type CreateProjectInput, + type ModelProvider, type ProjectManager, type ScaffoldHarnessInput, type ScaffoldRuntimeInput, @@ -65,7 +66,7 @@ const HARNESS_ONLY_FLAGS = [ "container", ] as const; -const ModelProviderFlagSchema = z.union([z.literal("Bedrock"), HarnessModelProviderSchema]); +const ModelProviderFlagSchema = z.enum([...HarnessModelProviderSchema.options, "anthropic"]); type ModelProviderFlag = z.infer; const HARNESS_DEFAULT_MODEL_IDS: Record = { @@ -116,7 +117,8 @@ export const createCreateProjectHandler = (config: CreateProjectHandlerConfig) = ), flag( "model-provider", - "model provider: bedrock, open_ai, gemini, or lite_llm for harnesses; Bedrock for runtime code", + "model provider: bedrock, open_ai, gemini, or lite_llm for harnesses; " + + "bedrock, anthropic, open_ai, or gemini for runtime code", ModelProviderFlagSchema.optional(), ), flag( @@ -411,18 +413,36 @@ export function resolveScaffoldHarnessInput(flags: HarnessPathFlagValues): Scaff return input; } -function resolveHarnessModelProvider(value: ModelProviderFlag | undefined): HarnessModelProvider { - return value === undefined || value === "Bedrock" ? "bedrock" : value; +// Runtimes and harnesses support different model sets and record them under +// different names in their spec configs, so the shared --model-provider flag is +// mapped to each domain here behind a consistent interface. +const MODEL_PROVIDERS: Record< + ModelProviderFlag, + { harness?: HarnessModelProvider; runtime?: ModelProvider } +> = { + bedrock: { harness: "bedrock", runtime: "Bedrock" }, + open_ai: { harness: "open_ai", runtime: "OpenAI" }, + gemini: { harness: "gemini", runtime: "Gemini" }, + lite_llm: { harness: "lite_llm", runtime: "LiteLLM" }, + anthropic: { runtime: "Anthropic" }, +}; + +function resolveHarnessModelProvider( + providerFlag: ModelProviderFlag | undefined, +): HarnessModelProvider { + if (providerFlag === undefined) return "bedrock"; + const provider = MODEL_PROVIDERS[providerFlag].harness; + if (provider === undefined) + throw new InputValidationError( + `the '${providerFlag}' model provider is not supported for harness projects`, + ); + return provider; } function resolveRuntimeModelProvider( - value: ModelProviderFlag | undefined, -): ScaffoldRuntimeInput["modelProvider"] | undefined { - if (value === undefined) return undefined; - if (value === "Bedrock" || value === "bedrock") return "Bedrock"; - throw new InputValidationError( - `runtime scaffolding only supports the Bedrock model provider; received '${value}'`, - ); + providerFlag: ModelProviderFlag | undefined, +): ModelProvider | undefined { + return providerFlag === undefined ? undefined : MODEL_PROVIDERS[providerFlag].runtime; } /** A --container value is either an ECR image URI or a local Dockerfile path. */ diff --git a/src/handlers/project/project.test.ts b/src/handlers/project/project.test.ts index 200f55ce6..20f0a4aeb 100644 --- a/src/handlers/project/project.test.ts +++ b/src/handlers/project/project.test.ts @@ -182,6 +182,13 @@ describe("project create", () => { }, ); + test("rejects a runtime-only provider on the harness path", async () => { + await inTempDirectory(); + await expect( + run(["create", "--name", "MyAgent", "--model-provider", "anthropic"]), + ).rejects.toThrow(/'anthropic' model provider is not supported for harness projects/); + }); + test("supports LiteLLM model configuration on the harness path", async () => { const directory = await inTempDirectory(); await run([ @@ -406,7 +413,7 @@ describe("project create", () => { "--build", "CodeZip", "--model-provider", - "Bedrock", + "bedrock", "--memory", "none", "--skip-install", @@ -424,20 +431,58 @@ describe("project create", () => { expect(await Bun.file(join(projectRoot, "app", "custom_agent", "main.py")).exists()).toBe(true); }); - test("rejects non-Bedrock model providers on the runtime path", async () => { + test("scaffolds a keyless LiteLLM runtime with no credential", async () => { const directory = await inTempDirectory(); - await expect( - run([ - "create", - "--name", - "MyProject", - "--template", - "agent-python-strands", - "--model-provider", - "open_ai", - ]), - ).rejects.toThrow(/runtime scaffolding only supports the Bedrock model provider/); - expect(existsSync(join(directory, "MyProject"))).toBe(false); + await run([ + "create", + "--name", + "MyProject", + "--template", + "agent-python-strands", + "--model-provider", + "lite_llm", + "--skip-install", + "--skip-git", + ]); + + const projectRoot = join(directory, "MyProject"); + const spec = await Bun.file(join(projectRoot, "agentcore", "agentcore.json")).json(); + expect(spec.runtimes).toHaveLength(1); + expect(spec.credentials ?? []).toEqual([]); + }); + + test.each<[string, string]>([ + ["anthropic", "agent_python_strandsAnthropicApiKey"], + ["open_ai", "agent_python_strandsOpenAIApiKey"], + ["gemini", "agent_python_strandsGeminiApiKey"], + ["lite_llm", "agent_python_strandsLiteLLMApiKey"], + ])("scaffolds a runtime with a %s API-key credential", async (provider, credentialName) => { + const directory = await inTempDirectory(); + const apiKeyPath = join(directory, "api-key.txt"); + await Bun.write(apiKeyPath, "test-api-key"); + + await run([ + "create", + "--name", + "MyProject", + "--template", + "agent-python-strands", + "--model-provider", + provider, + "--api-key", + `file://${apiKeyPath}`, + "--skip-install", + "--skip-git", + ]); + + const projectRoot = join(directory, "MyProject"); + const spec = await Bun.file(join(projectRoot, "agentcore", "agentcore.json")).json(); + expect(spec.credentials).toContainEqual({ + authorizerType: "ApiKeyCredentialProvider", + name: credentialName, + }); + const envLocal = await Bun.file(join(projectRoot, "agentcore", ".env.local")).text(); + expect(envLocal).toContain("test-api-key"); }); test("scaffolds a Container agent from the strands template", async () => { @@ -583,7 +628,7 @@ describe("project create", () => { "--framework", "strands", "--model-provider", - "Bedrock", + "bedrock", ...memoryFlags, "--skip-install", "--skip-git", @@ -619,7 +664,7 @@ describe("project create", () => { "--framework", "none", "--model-provider", - "Bedrock", + "bedrock", "--memory", "none", "--skip-install", @@ -652,7 +697,7 @@ describe("project create", () => { "--framework", "strands", "--model-provider", - "Bedrock", + "bedrock", "--memory", "none", "--skip-install", @@ -692,7 +737,7 @@ describe("project create", () => { "--framework", "none", "--model-provider", - "Bedrock", + "bedrock", "--memory", memoryShortcut, "--skip-install", @@ -726,7 +771,7 @@ describe("project create", () => { "--framework", "none", "--model-provider", - "Bedrock", + "bedrock", "--memory", "none", "--skip-install", diff --git a/src/handlers/project/shortcuts.ts b/src/handlers/project/shortcuts.ts index 72a6e08da..f1ec2d4a0 100644 --- a/src/handlers/project/shortcuts.ts +++ b/src/handlers/project/shortcuts.ts @@ -83,7 +83,6 @@ export const RUNTIME_TEMPLATE_SHORTCUTS = { language: "Python", framework: "none", protocol: "MCP", - modelProvider: "Bedrock", memory: "none", runtimeVersion: "PYTHON_3_14", }, diff --git a/src/handlers/project/types.ts b/src/handlers/project/types.ts index dc2dd279d..d2649a328 100644 --- a/src/handlers/project/types.ts +++ b/src/handlers/project/types.ts @@ -40,6 +40,27 @@ export type ManagedEvaluatorScaffoldInput = { timeoutSeconds?: number; }; +/** Model providers the scaffolded runtime code can target. */ +export const MODEL_PROVIDERS = ["Bedrock", "Anthropic", "OpenAI", "Gemini", "LiteLLM"] as const; +export type ModelProvider = (typeof MODEL_PROVIDERS)[number]; + +const MODEL_PROVIDER_ALIASES: Record = { + bedrock: "Bedrock", + anthropic: "Anthropic", + openai: "OpenAI", + open_ai: "OpenAI", + gemini: "Gemini", + litellm: "LiteLLM", + lite_llm: "LiteLLM", +}; + +/** Parses a provider name case-insensitively (e.g. `anthropic`), normalizing to canonical casing. */ +export const ModelProviderSchema = z.preprocess( + (value) => + typeof value === "string" ? (MODEL_PROVIDER_ALIASES[value.toLowerCase()] ?? value) : value, + z.enum(MODEL_PROVIDERS), +); + /** Set of arguments needed to scaffold a new Runtime-based agent. */ export const ScaffoldRuntimeInputSchema = z .object({ @@ -48,14 +69,31 @@ export const ScaffoldRuntimeInputSchema = z language: z.enum(["Python", "TypeScript"]), framework: z.enum(["strands", "none"]), protocol: ProtocolModeSchema.optional(), - modelProvider: z.enum(["Bedrock"]), + modelProvider: ModelProviderSchema.optional(), apiKey: z.string().min(1).optional(), memory: MemorySchema.optional(), runtimeVersion: RuntimeVersionSchema.optional(), }) - .refine(({ modelProvider, apiKey }) => !(modelProvider === "Bedrock" && apiKey !== undefined), { - message: "API keys are not compatible with Bedrock model providers", - path: ["apiKey"], + .superRefine(({ modelProvider, apiKey }, ctx) => { + // LiteLLM routes to any provider (Bedrock by default), so its key is optional; + // the other non-Bedrock providers always call their own API and require one. + const requiresApiKey = + modelProvider !== undefined && modelProvider !== "Bedrock" && modelProvider !== "LiteLLM"; + const allowsApiKey = requiresApiKey || modelProvider === "LiteLLM"; + if (apiKey !== undefined && !allowsApiKey) { + ctx.addIssue({ + code: "custom", + message: "API keys are not compatible with Bedrock model providers", + path: ["apiKey"], + }); + } + if (apiKey === undefined && requiresApiKey) { + ctx.addIssue({ + code: "custom", + message: `an API key is required for the ${modelProvider} model provider`, + path: ["apiKey"], + }); + } }) .superRefine(({ build, runtimeVersion }, ctx) => { if (build === "CodeZip" && runtimeVersion === undefined) {