diff --git a/src/assets/templates/strands-http-python/memory/__init__.py b/src/assets/templates/strands-http-python/memory/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/src/assets/templates/strands-http-python/memory/session.py b/src/assets/templates/strands-http-python/memory/session.py new file mode 100644 index 000000000..20e105674 --- /dev/null +++ b/src/assets/templates/strands-http-python/memory/session.py @@ -0,0 +1,47 @@ +import os +import uuid +from typing import Optional + +from bedrock_agentcore.memory.integrations.strands.config import AgentCoreMemoryConfig{{#if memoryStrategies.length}}, RetrievalConfig{{/if}} +from bedrock_agentcore.memory.integrations.strands.session_manager import AgentCoreMemorySessionManager + +MEMORY_ID = os.getenv("{{memoryEnvVarName}}") +REGION = os.getenv("AWS_REGION") + + +def get_memory_session_manager( + session_id: Optional[str], actor_id: str +) -> Optional[AgentCoreMemorySessionManager]: + if not MEMORY_ID: + return None + + session_id = session_id or uuid.uuid4().hex + +{{#if memoryStrategies.length}} + retrieval_config = { +{{#if (includes memoryStrategies "SEMANTIC")}} + f"/users/{actor_id}/facts": RetrievalConfig(top_k=3, relevance_score=0.5), +{{/if}} +{{#if (includes memoryStrategies "USER_PREFERENCE")}} + f"/users/{actor_id}/preferences": RetrievalConfig(top_k=3, relevance_score=0.5), +{{/if}} +{{#if (includes memoryStrategies "EPISODIC")}} + f"/episodes/{actor_id}/{session_id}": RetrievalConfig(top_k=5, relevance_score=0.5), +{{/if}} +{{#if (includes memoryStrategies "SUMMARIZATION")}} + f"/summaries/{actor_id}": RetrievalConfig(top_k=3, relevance_score=0.5), +{{/if}} + } +{{/if}} + + return AgentCoreMemorySessionManager( + AgentCoreMemoryConfig( + memory_id=MEMORY_ID, + session_id=session_id, + actor_id=actor_id, +{{#if memoryStrategies.length}} + retrieval_config=retrieval_config, +{{/if}} + ), + REGION, + ) diff --git a/src/core/project/__snapshots__/manager.test.ts.snap b/src/core/project/__snapshots__/manager.test.ts.snap index 6fd45fcf4..2a538c1bb 100644 --- a/src/core/project/__snapshots__/manager.test.ts.snap +++ b/src/core/project/__snapshots__/manager.test.ts.snap @@ -46,12 +46,49 @@ exports[`FsProjectManager.create snapshots the Strands project manifest and runt "app/strands_agent/main.py", "app/strands_agent/mcp_client/__init__.py", "app/strands_agent/mcp_client/client.py", + "app/strands_agent/memory/__init__.py", + "app/strands_agent/memory/session.py", "app/strands_agent/model/__init__.py", "app/strands_agent/model/load.py", "app/strands_agent/model/mantle_compat.py", "app/strands_agent/pyproject.toml", "app/strands_agent/skills/fetcher.py", ], + "memories": [ + { + "eventExpiryDuration": 30, + "name": "strands_agentMemory", + "strategies": [ + { + "namespaceTemplates": [ + "/users/{actorId}/facts", + ], + "type": "SEMANTIC", + }, + { + "namespaceTemplates": [ + "/users/{actorId}/preferences", + ], + "type": "USER_PREFERENCE", + }, + { + "namespaceTemplates": [ + "/summaries/{actorId}/{sessionId}", + ], + "type": "SUMMARIZATION", + }, + { + "namespaceTemplates": [ + "/episodes/{actorId}/{sessionId}", + ], + "reflectionNamespaceTemplates": [ + "/episodes/{actorId}", + ], + "type": "EPISODIC", + }, + ], + }, + ], "runtimes": [ { "build": "CodeZip", diff --git a/src/core/project/manager.test.ts b/src/core/project/manager.test.ts index 620819493..89c435159 100644 --- a/src/core/project/manager.test.ts +++ b/src/core/project/manager.test.ts @@ -6,7 +6,7 @@ import { DeserializationError, ProjectStateError } from "../../errors/errors"; import type { AwsDeploymentTarget } from "../../projectSchemas/aws-targets"; import { ProjectSpecSchema } from "../../projectSchemas/project"; import { FsProjectManager } from "./manager"; -import { RUNTIME_TEMPLATE_SHORTCUTS } from "../../handlers/project/shortcuts"; +import { resolveRuntimeTemplateShortcut } from "../../handlers/project/shortcuts"; import { type CreateProjectInput, type DeployResult, @@ -17,9 +17,9 @@ import { import { createSilentLogger } from "../../testing"; import type { DeployBackendInput, ProjectBackend } from "./backends/types"; -const HELLO_WORLD_PYTHON = RUNTIME_TEMPLATE_SHORTCUTS["hello-world-python"]; -const HELLO_WORLD_PYTHON_CONTAINER = RUNTIME_TEMPLATE_SHORTCUTS["hello-world-python-container"]; -const STRANDS_PYTHON = RUNTIME_TEMPLATE_SHORTCUTS["strands-python"]; +const HELLO_WORLD_PYTHON = resolveRuntimeTemplateShortcut("hello-world-python"); +const HELLO_WORLD_PYTHON_CONTAINER = resolveRuntimeTemplateShortcut("hello-world-python-container"); +const STRANDS_PYTHON = resolveRuntimeTemplateShortcut("strands-python"); const originalCwd = process.cwd(); const tempDirectories: string[] = []; @@ -102,6 +102,7 @@ describe("FsProjectManager.create", () => { expect({ manifest: await projectManifest(projectRoot), runtimes: spec.runtimes, + memories: spec.memories, }).toMatchSnapshot(); }); diff --git a/src/core/project/templates/fsTree.test.ts b/src/core/project/templates/fsTree.test.ts index 1ad93cc74..c06f3db1f 100644 --- a/src/core/project/templates/fsTree.test.ts +++ b/src/core/project/templates/fsTree.test.ts @@ -72,7 +72,11 @@ describe("FsTreeNode.fromAssetSource", () => { }, }; - const tree = await FsTreeNode.fromAssetSource(source, "template", "root"); + const tree = await FsTreeNode.fromAssetSource( + { assetSource: source }, + { assetDir: "template" }, + { rootDirName: "root" }, + ); expect(tree.name).toBe("root"); expect(tree.children.map((node) => node.name)).toEqual(["README.md", "src", ".gitignore"]); @@ -84,6 +88,29 @@ describe("FsTreeNode.fromAssetSource", () => { expect(await tree.children[2]?.bytes?.()).toBe("contents:template/gitignore.template"); }); + test("transforms content and filters files and directories", async () => { + const source: AssetSource = { + async list() { + return ["template/keep.txt", "template/skip.txt", "template/optional/nested.txt"]; + }, + async read(assetPath) { + return `contents:${assetPath}`; + }, + }; + + const tree = await FsTreeNode.fromAssetSource( + { assetSource: source }, + { assetDir: "template" }, + { + transformContent: (content) => content.toUpperCase(), + filter: (name, isDir) => name !== "skip.txt" && !(isDir && name === "optional"), + }, + ); + + expect(tree.children.map(({ name }) => name)).toEqual(["keep.txt"]); + expect(await tree.children[0]?.bytes?.()).toBe("CONTENTS:TEMPLATE/KEEP.TXT"); + }); + test("strips .template suffix from non-ignore files", async () => { const source: AssetSource = { async list() { @@ -94,7 +121,11 @@ describe("FsTreeNode.fromAssetSource", () => { }, }; - const tree = await FsTreeNode.fromAssetSource(source, "template", "root"); + const tree = await FsTreeNode.fromAssetSource( + { assetSource: source }, + { assetDir: "template" }, + { rootDirName: "root" }, + ); expect(tree.children.map((node) => node.name)).toEqual(["Dockerfile", ".dockerignore"]); }); diff --git a/src/core/project/templates/fsTree.ts b/src/core/project/templates/fsTree.ts index ec9171d42..c1281ebc3 100644 --- a/src/core/project/templates/fsTree.ts +++ b/src/core/project/templates/fsTree.ts @@ -71,18 +71,30 @@ export class FsTreeNode { } /** - * Expands the flat asset listing under assetDir into a nested tree of nodes. + * Builds a file tree from assets under `input.assetDir`. + * + * @param config - Asset source configuration. + * @param input - Asset directory to load. + * @param options - Optional root name, lazy content transform, and descendant filter. Rejecting a directory omits its subtree. */ static async fromAssetSource( - src: AssetSource, - assetDir: string, - rootDirName?: string, - transform?: (content: string) => string, + config: { assetSource: AssetSource }, + input: { assetDir: string }, + options?: { + rootDirName?: string; + transformContent?: (content: string) => string; + filter?: (name: string, isDir: boolean) => boolean; + }, ): Promise { - const paths = await src.list(assetDir); + const { assetSource } = config; + const { assetDir } = input; + const rootDirName = options?.rootDirName; + const transformContent = options?.transformContent; + const filter = options?.filter; + const paths = await assetSource.list(assetDir); const root = FsTreeNode.createDirectory(rootDirName ?? assetDir, []); - for (const assetPath of paths) { + assetPaths: for (const assetPath of paths) { const relative = assetPath.slice(assetDir.length + 1); const segments = relative.split("/"); if (segments.some((s) => s === "" || s === "." || s === "..")) { @@ -92,25 +104,31 @@ export class FsTreeNode { } let parent = root; - segments.forEach((segment, index) => { - if (index === segments.length - 1) { + for (const [index, segment] of segments.entries()) { + const isDir = index < segments.length - 1; + const name = isDir ? segment : renderName(segment); + // if the segment of a path rejects, reject the rest of the path so we jump to top-loop via assetPaths label. + if (filter && !filter(name, isDir)) continue assetPaths; + + if (!isDir) { parent.children.push( - FsTreeNode.createFile(renderName(segment), async () => { - const raw = await src.read(assetPath); - return transform ? transform(raw) : raw; + FsTreeNode.createFile(name, async () => { + const raw = await assetSource.read(assetPath); + return transformContent ? transformContent(raw) : raw; }), ); - return; + continue; } - let child = parent.children.find((n): n is FsTreeNode => n.isDir && n.name === segment); + let child = parent.children.find( + (node): node is FsTreeNode => node.isDir && node.name === name, + ); if (!child) { - child = FsTreeNode.createDirectory(segment, []); + child = FsTreeNode.createDirectory(name, []); parent.children.push(child); } - parent = child; - }); + } } return root; diff --git a/src/core/project/templates/project.ts b/src/core/project/templates/project.ts index ff0a1ecaa..373fcd43b 100644 --- a/src/core/project/templates/project.ts +++ b/src/core/project/templates/project.ts @@ -35,7 +35,7 @@ export async function createProjectTree( config.assetSource.read("templates/shared/gitignore.template"), ), FsTreeNode.createDirectory("agentcore", [ - await FsTreeNode.fromAssetSource(config.assetSource, "cdk"), + await FsTreeNode.fromAssetSource({ assetSource: config.assetSource }, { assetDir: "cdk" }), FsTreeNode.createFile("agentcore.json", async () => json({ name: input.projectName, diff --git a/src/core/project/templates/runtime.ts b/src/core/project/templates/runtime.ts index f1db624f5..6a15566d1 100644 --- a/src/core/project/templates/runtime.ts +++ b/src/core/project/templates/runtime.ts @@ -60,12 +60,17 @@ const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: Templa [buildResolverKey("none", "Python")]: async (input: RuntimeResourceConfig) => { if (input.protocol !== undefined && input.protocol !== "HTTP") throw new InputValidationError(`hello-world-python only supports HTTP protocol`); + if (input.scaffoldRuntimeInput.memory !== undefined) + throw new InputValidationError(`memory is not supported with the hello-world template`); const tree = await FsTreeNode.fromAssetSource( - assetSource, - input.scaffoldRuntimeInput.build === "Container" - ? "templates/hello-world-python-container" - : "templates/hello-world-python", - input.name, + { assetSource }, + { + assetDir: + input.scaffoldRuntimeInput.build === "Container" + ? "templates/hello-world-python-container" + : "templates/hello-world-python", + }, + { rootDirName: input.name }, ); return { tree, spec: { runtimes: [buildRuntimeSpec(input)] } }; }, @@ -90,10 +95,14 @@ const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: Templa ? [{ mountPath: configuration.s3FilesAccessPoint.mountPath }] : [], ); + const memory = input.scaffoldRuntimeInput.memory; const context = { name: toPythonPackageName(input.name), modelProvider: input.scaffoldRuntimeInput.modelProvider, - hasMemory: input.scaffoldRuntimeInput.memory !== "none", + 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, hasGateway: false, hasPayment: false, @@ -108,14 +117,20 @@ const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: Templa hasConfigBundle: false, }; const tree = await FsTreeNode.fromAssetSource( - assetSource, - "templates/strands-http-python", - input.name, - (raw) => templateRenderer.render(raw, context), + { assetSource }, + { assetDir: "templates/strands-http-python" }, + { + rootDirName: input.name, + transformContent: (raw) => templateRenderer.render(raw, context), + filter: (name, isDir) => memory !== undefined || !isDir || name !== "memory", + }, ); return { tree, - spec: { runtimes: [{ ...buildRuntimeSpec(input), protocol: "HTTP" as const }] }, + spec: { + runtimes: [{ ...buildRuntimeSpec(input), protocol: "HTTP" as const }], + ...(memory && { memories: [memory] }), + }, }; }, }); diff --git a/src/handlers/project/add/runtime/index.test.ts b/src/handlers/project/add/runtime/index.test.ts index 6b6a29f5b..e5c2f047f 100644 --- a/src/handlers/project/add/runtime/index.test.ts +++ b/src/handlers/project/add/runtime/index.test.ts @@ -321,6 +321,50 @@ describe("project add runtime", () => { expect(runtime.runtimeVersion).toBe(isContainer ? undefined : "PYTHON_3_14"); }); + test.each([ + ["default", [], ["SEMANTIC", "USER_PREFERENCE", "SUMMARIZATION", "EPISODIC"]], + ["none", ["--memory", "none"], []], + ["short", ["--memory", "shortTerm"], []], + [ + "longAndShortTerm", + ["--memory", "longAndShortTerm"], + ["SEMANTIC", "USER_PREFERENCE", "SUMMARIZATION", "EPISODIC"], + ], + ])("custom strands %s memory", async (_label, memoryFlags, expectedStrategies) => { + const projectRoot = await inProject(); + await run([ + "add", + "runtime", + "--name", + "my_agent", + "--build", + "CodeZip", + "--language", + "Python", + "--framework", + "strands", + "--model-provider", + "Bedrock", + ...memoryFlags, + ]); + + const spec = await Bun.file(join(projectRoot, "agentcore", "agentcore.json")).json(); + const memory = spec.memories.find( + (candidate: { name: string }) => candidate.name === "my_agentMemory", + ); + + if (memoryFlags.length > 1 && memoryFlags[1] === "none") { + expect(memory).toBeUndefined(); + return; + } + + expect(memory).toMatchObject({ + name: "my_agentMemory", + eventExpiryDuration: 30, + }); + expect(memory.strategies.map(({ type }: { type: string }) => type)).toEqual(expectedStrategies); + }); + test.each<[string, string[]]>([ ["missing --name", ["--template", "hello-world-python"]], [ @@ -354,6 +398,41 @@ describe("project add runtime", () => { "hello-world-python only supports HTTP", ["--name", "my_agent", "--template", "hello-world-python", "--protocol", "MCP"], ], + [ + "--memory shortTerm is not supported with --framework none", + [ + "--name", + "my_agent", + "--build", + "CodeZip", + "--language", + "Python", + "--framework", + "none", + "--model-provider", + "Bedrock", + "--memory", + "shortTerm", + ], + ], + [ + "--memory longAndShortTerm is not supported with --framework none", + [ + "--name", + "my_agent", + "--build", + "CodeZip", + "--language", + "Python", + "--framework", + "none", + "--model-provider", + "Bedrock", + "--memory", + "longAndShortTerm", + ], + ], + ["runtime names are limited in length", ["--name", "x".repeat(43)]], ])("%s", async (_label, flags) => { await inProject(); await expect(run(["add", "runtime", ...flags])).rejects.toBeInstanceOf(InputValidationError); diff --git a/src/handlers/project/add/runtime/index.ts b/src/handlers/project/add/runtime/index.ts index c4ece055b..4a753a59a 100644 --- a/src/handlers/project/add/runtime/index.ts +++ b/src/handlers/project/add/runtime/index.ts @@ -8,11 +8,12 @@ import { RuntimeAuthorizerTypeSchema } from "../../../../projectSchemas/auth"; import { NetworkModeSchema, ProtocolModeSchema } from "../../../../projectSchemas/constants"; import { SourceResolver } from "../../../../io"; import { + MEMORY_SHORTCUT_NAMES, + MEMORY_SHORTCUTS, RUNTIME_TEMPLATE_SHORTCUT_NAMES, - RUNTIME_TEMPLATE_SHORTCUTS, resolveRuntimeTemplateShortcut, } from "../../shortcuts"; -import { ScaffoldRuntimeInputSchema } from "../../types"; +import { ScaffoldRuntimeInputSchema, type ScaffoldRuntimeInput } from "../../types"; import { RuntimeResourceConfigSchema } from "./types"; export const createAddRuntimeHandler = (config: AddProjectResourceConfig) => @@ -20,7 +21,7 @@ export const createAddRuntimeHandler = (config: AddProjectResourceConfig) => name: "runtime", description: "adds a runtime to the current project", flags: [ - flag("name", "the name of the runtime", z.string().optional()), + flag("name", "the name of the runtime", z.string().max(42).optional()), flag("description", "an optional description of the runtime", z.string().optional()), flag( "template", @@ -49,7 +50,11 @@ export const createAddRuntimeHandler = (config: AddProjectResourceConfig) => z.string().optional(), { sensitive: true }, ), - flag("memory", "memory option for the scaffolded runtime", z.enum(["none"]).optional()), + flag( + "memory", + "memory option for the scaffolded runtime", + z.enum(MEMORY_SHORTCUT_NAMES).optional(), + ), flag( "role-arn", "IAM role ARN that provides permissions for the runtime", @@ -121,32 +126,29 @@ export const createAddRuntimeHandler = (config: AddProjectResourceConfig) => const source = new SourceResolver({ stdin: config.io.stdin }); const apiKey = await source.resolveSecret("api-key", flags["api-key"]); - const scaffoldRuntimeInput = isTemplate + const runtimeName = flags.name; + const defaultMemory = flags.framework === "strands" ? "longAndShortTerm" : "none"; + const scaffoldRuntimeInput: ScaffoldRuntimeInput = isTemplate ? resolveRuntimeTemplateShortcut(flags.template!, { runtimeName: flags.name, - ...(flags.build !== undefined && { - build: flags.build, - runtimeVersion: flags.build === "CodeZip" ? "PYTHON_3_14" : undefined, - }), - ...(flags["model-provider"] !== undefined && { - modelProvider: flags["model-provider"], - }), - ...(apiKey !== undefined && { apiKey }), - ...(flags.memory !== undefined && { memory: flags.memory }), + build: flags.build, + modelProvider: flags["model-provider"], + apiKey, + memory: flags.memory, }) : isCustom ? parseScaffoldRuntimeInput({ - runtimeName: flags.name, + runtimeName, build: flags.build, language: flags.language, framework: flags.framework, modelProvider: flags["model-provider"], apiKey, - memory: flags.memory, + memory: MEMORY_SHORTCUTS[flags.memory ?? defaultMemory](runtimeName), entrypoint: "main.py", runtimeVersion: flags.build === "CodeZip" ? "PYTHON_3_14" : undefined, }) - : RUNTIME_TEMPLATE_SHORTCUTS["hello-world-python"]; + : resolveRuntimeTemplateShortcut("hello-world-python", { runtimeName: flags.name }); const inputEnvironmentVariables = parseJsonFlag>( "environment-variables", @@ -200,7 +202,7 @@ function toEnvironmentVariables(envVars: Record | undefined): En return envVars ? Object.entries(envVars).map(([name, value]) => ({ name, value })) : []; } -function parseScaffoldRuntimeInput(input: Record) { +function parseScaffoldRuntimeInput(input: Partial) { const result = ScaffoldRuntimeInputSchema.safeParse(input); if (!result.success) throw new InputValidationError(z.prettifyError(result.error)); return result.data; diff --git a/src/handlers/project/create/index.ts b/src/handlers/project/create/index.ts index 29d537f72..69441eff5 100644 --- a/src/handlers/project/create/index.ts +++ b/src/handlers/project/create/index.ts @@ -2,11 +2,17 @@ import z from "zod"; import { createHandler, flag } from "../../../router"; import { SourceResolver, type AppIO } from "../../../io"; import { + MEMORY_SHORTCUT_NAMES, + MEMORY_SHORTCUTS, RUNTIME_TEMPLATE_SHORTCUT_NAMES, - RUNTIME_TEMPLATE_SHORTCUTS, resolveRuntimeTemplateShortcut, } from "../shortcuts"; -import { ScaffoldRuntimeInputSchema, type CreateProjectInput, type ProjectManager } from "../types"; +import { + ScaffoldRuntimeInputSchema, + type CreateProjectInput, + type ProjectManager, + type ScaffoldRuntimeInput, +} from "../types"; import { ProjectNameSchema } from "../../../projectSchemas/project"; import { InputValidationError } from "../../../errors"; @@ -52,8 +58,12 @@ export const createCreateProjectHandler = (config: CreateProjectHandlerConfig) = z.string().optional(), { sensitive: true }, ), - flag("memory", "memory option for the scaffolded runtime", z.enum(["none"]).optional()), - flag("runtime-name", "name of the scaffolded runtime", z.string().optional()), + flag( + "memory", + "memory option for the scaffolded runtime", + z.enum(MEMORY_SHORTCUT_NAMES).optional(), + ), + flag("runtime-name", "name of the scaffolded runtime", z.string().max(42).optional()), flag( "skip-install", "skip installing dependencies (npm install, uv sync)", @@ -68,8 +78,8 @@ export const createCreateProjectHandler = (config: CreateProjectHandlerConfig) = "framework", "model-provider", "api-key", - "memory", "runtime-name", + "memory", ] as const; const presentScaffoldingFlags = scaffoldingFlags.filter((f) => flags[f] !== undefined); @@ -86,34 +96,30 @@ export const createCreateProjectHandler = (config: CreateProjectHandlerConfig) = const source = new SourceResolver({ stdin: config.io.stdin }); const apiKey = await source.resolveSecret("api-key", flags["api-key"]); - const scaffoldRuntimeInput = isTemplate + const runtimeName = flags["runtime-name"] ?? flags["name"]; + const defaultMemory = flags["framework"] === "strands" ? "longAndShortTerm" : "none"; + + const scaffoldRuntimeInput: ScaffoldRuntimeInput = isTemplate ? resolveRuntimeTemplateShortcut(flags["template"]!, { - ...(flags["runtime-name"] !== undefined && { - runtimeName: flags["runtime-name"], - }), - ...(flags["build"] !== undefined && { - build: flags["build"], - runtimeVersion: flags["build"] === "CodeZip" ? "PYTHON_3_14" : undefined, - }), - ...(flags["model-provider"] !== undefined && { - modelProvider: flags["model-provider"], - }), - ...(apiKey !== undefined && { apiKey }), - ...(flags["memory"] !== undefined && { memory: flags["memory"] }), + runtimeName: flags["runtime-name"], + build: flags["build"], + modelProvider: flags["model-provider"], + apiKey, + memory: flags["memory"], }) : isCustom ? parseScaffoldRuntimeInput({ - runtimeName: flags["runtime-name"] ?? flags["name"], + runtimeName, build: flags["build"], language: flags["language"], framework: flags["framework"], modelProvider: flags["model-provider"], apiKey, - memory: flags["memory"], + memory: MEMORY_SHORTCUTS[flags["memory"] ?? defaultMemory](runtimeName), entrypoint: "main.py", runtimeVersion: flags["build"] === "CodeZip" ? "PYTHON_3_14" : undefined, }) - : RUNTIME_TEMPLATE_SHORTCUTS["hello-world-python"]; + : resolveRuntimeTemplateShortcut("hello-world-python"); const createInput: CreateProjectInput = { name: flags["name"], @@ -130,7 +136,7 @@ export const createCreateProjectHandler = (config: CreateProjectHandlerConfig) = }, }); -function parseScaffoldRuntimeInput(input: Record) { +function parseScaffoldRuntimeInput(input: Partial) { const result = ScaffoldRuntimeInputSchema.safeParse(input); if (!result.success) throw new InputValidationError(z.prettifyError(result.error)); return result.data; diff --git a/src/handlers/project/project.test.ts b/src/handlers/project/project.test.ts index 50ae1d09b..742a4965c 100644 --- a/src/handlers/project/project.test.ts +++ b/src/handlers/project/project.test.ts @@ -163,6 +163,51 @@ describe("project create", () => { expect(await Bun.file(join(projectRoot, "app", "custom_agent", "main.py")).exists()).toBe(true); }); + test.each([ + ["default", [], ["SEMANTIC", "USER_PREFERENCE", "SUMMARIZATION", "EPISODIC"]], + ["none", ["--memory", "none"], []], + ["short", ["--memory", "shortTerm"], []], + [ + "shortAndLongTerm", + ["--memory", "longAndShortTerm"], + ["SEMANTIC", "USER_PREFERENCE", "SUMMARIZATION", "EPISODIC"], + ], + ])("custom strands %s memory", async (_label, memoryFlags, expectedStrategies) => { + const directory = await inTempDirectory(); + await run([ + "create", + "--name", + "MyAgent", + "--build", + "CodeZip", + "--language", + "Python", + "--framework", + "strands", + "--model-provider", + "Bedrock", + ...memoryFlags, + "--skip-install", + "--skip-git", + ]); + + const projectRoot = join(directory, "MyAgent"); + const spec = await Bun.file(join(projectRoot, "agentcore", "agentcore.json")).json(); + const memories = spec.memories ?? []; + const memory = memories[0]; + + if (memoryFlags.length > 1 && memoryFlags[1] === "none") { + expect(memories).toEqual([]); + return; + } + + expect(memory).toMatchObject({ + name: "MyAgentMemory", + eventExpiryDuration: 30, + }); + expect(memory.strategies.map(({ type }: { type: string }) => type)).toEqual(expectedStrategies); + }); + test("scaffolds from explicit custom flags", async () => { const directory = await inTempDirectory(); await run([ @@ -196,32 +241,67 @@ describe("project create", () => { ]); }); - test("rejects an invalid --runtime-name before scaffolding", async () => { - const directory = await inTempDirectory(); - await expect( - run([ - "create", - "--name", - "MyProject", - "--runtime-name", - "../MyAgent", - "--build", - "CodeZip", - "--language", - "Python", - "--framework", - "none", - "--model-provider", - "Bedrock", - "--memory", - "none", - "--skip-install", - "--skip-git", - ]), - ).rejects.toThrow(/Must begin with a letter/); + test.each(["shortTerm", "longAndShortTerm"] as const)( + "rejects --memory %s with --framework none", + async (memoryShortcut) => { + await inTempDirectory(); + await expect( + run([ + "create", + "--name", + "MyAgent", + "--build", + "CodeZip", + "--language", + "Python", + "--framework", + "none", + "--model-provider", + "Bedrock", + "--memory", + memoryShortcut, + "--skip-install", + "--skip-git", + ]), + ).rejects.toBeInstanceOf(InputValidationError); + }, + ); - expect(existsSync(join(directory, "MyProject"))).toBe(false); - }); + test.each([ + ["path traversal", "../MyAgent", /Must begin with a letter/], + ["starts with a digit", "1Agent", /Must begin with a letter/], + ["contains a hyphen", "my-agent", /Must begin with a letter/], + ["contains a space", "my agent", /Must begin with a letter/], + ["exceeds 42 chars", "a".repeat(43), /<=42 characters/], + ])( + "rejects an invalid --runtime-name before scaffolding (%s)", + async (_label, runtimeName, expectedError) => { + const directory = await inTempDirectory(); + await expect( + run([ + "create", + "--name", + "MyProject", + "--runtime-name", + runtimeName, + "--build", + "CodeZip", + "--language", + "Python", + "--framework", + "none", + "--model-provider", + "Bedrock", + "--memory", + "none", + "--skip-install", + "--skip-git", + ]), + ).rejects.toThrow(expectedError); + + expect(existsSync(join(directory, "MyProject"))).toBe(false); + }, + ); test("rejects an incompatible API-key template override before scaffolding", async () => { const directory = await inTempDirectory(); diff --git a/src/handlers/project/shortcuts.ts b/src/handlers/project/shortcuts.ts index 86d62df13..81e81a77b 100644 --- a/src/handlers/project/shortcuts.ts +++ b/src/handlers/project/shortcuts.ts @@ -1,7 +1,45 @@ import z from "zod"; +import { + DEFAULT_EPISODIC_REFLECTION_NAMESPACE_TEMPLATES, + DEFAULT_STRATEGY_NAMESPACE_TEMPLATES, + type Memory, +} from "../../projectSchemas/memory"; import { InputValidationError } from "../../errors"; import { ScaffoldRuntimeInputSchema, type ScaffoldRuntimeInput } from "./types"; +export const MEMORY_SHORTCUTS = { + none: (_runtimeName: string) => undefined, + shortTerm: (runtimeName: string): Memory => ({ + name: `${runtimeName}Memory`, + eventExpiryDuration: 30, + strategies: [], + }), + longAndShortTerm: (runtimeName: string): Memory => ({ + name: `${runtimeName}Memory`, + eventExpiryDuration: 30, + strategies: (["SEMANTIC", "USER_PREFERENCE", "SUMMARIZATION", "EPISODIC"] as const).map( + (type) => ({ + type, + namespaceTemplates: DEFAULT_STRATEGY_NAMESPACE_TEMPLATES[type], + ...(type === "EPISODIC" && { + reflectionNamespaceTemplates: DEFAULT_EPISODIC_REFLECTION_NAMESPACE_TEMPLATES, + }), + }), + ), + }), +} satisfies Record Memory | undefined>; + +export type MemoryShortcutName = keyof typeof MEMORY_SHORTCUTS; + +export const MEMORY_SHORTCUT_NAMES = Object.keys(MEMORY_SHORTCUTS) as unknown as readonly [ + MemoryShortcutName, + ...MemoryShortcutName[], +]; + +type RuntimeTemplateShortcut = Omit & { + memory: MemoryShortcutName; +}; + export const RUNTIME_TEMPLATE_SHORTCUTS = { "hello-world-python": { runtimeName: "hello_world", @@ -28,11 +66,11 @@ export const RUNTIME_TEMPLATE_SHORTCUTS = { language: "Python", framework: "strands", modelProvider: "Bedrock", - memory: "none", + memory: "longAndShortTerm", entrypoint: "main.py", runtimeVersion: "PYTHON_3_14", }, -} as const satisfies Record; +} as const satisfies Record; export type RuntimeTemplateShortcutName = keyof typeof RUNTIME_TEMPLATE_SHORTCUTS; @@ -40,18 +78,35 @@ export const RUNTIME_TEMPLATE_SHORTCUT_NAMES = Object.keys( RUNTIME_TEMPLATE_SHORTCUTS, ) as unknown as readonly [RuntimeTemplateShortcutName, ...RuntimeTemplateShortcutName[]]; -type RuntimeTemplateOverrides = Partial< - Pick< - ScaffoldRuntimeInput, - "runtimeName" | "build" | "modelProvider" | "apiKey" | "memory" | "runtimeVersion" - > ->; +type RuntimeTemplateOverrides = { + runtimeName?: string; + build?: ScaffoldRuntimeInput["build"]; + modelProvider?: ScaffoldRuntimeInput["modelProvider"]; + apiKey?: string; + memory?: MemoryShortcutName; +}; export function resolveRuntimeTemplateShortcut( name: RuntimeTemplateShortcutName, - overrides: RuntimeTemplateOverrides, + overrides?: RuntimeTemplateOverrides, ): ScaffoldRuntimeInput { - const input = { ...RUNTIME_TEMPLATE_SHORTCUTS[name], ...overrides }; + const template = RUNTIME_TEMPLATE_SHORTCUTS[name]; + const runtimeName = overrides?.runtimeName ?? template.runtimeName; + const build = overrides?.build ?? template.build; + const memoryShortcutName = overrides?.memory ?? template.memory; + const memory = MEMORY_SHORTCUTS[memoryShortcutName](runtimeName); + + const input = { + runtimeName, + build, + language: template.language, + framework: template.framework, + modelProvider: overrides?.modelProvider ?? template.modelProvider, + ...(overrides?.apiKey !== undefined && { apiKey: overrides.apiKey }), + ...(memory && { memory }), + entrypoint: template.entrypoint, + runtimeVersion: build === "CodeZip" ? "PYTHON_3_14" : undefined, + }; const result = ScaffoldRuntimeInputSchema.safeParse(input); if (!result.success) throw new InputValidationError(z.prettifyError(result.error)); diff --git a/src/handlers/project/types.ts b/src/handlers/project/types.ts index 4b12cb84e..19b3a8ef4 100644 --- a/src/handlers/project/types.ts +++ b/src/handlers/project/types.ts @@ -1,7 +1,7 @@ import { HarnessSpecSchema } from "../../projectSchemas/harness"; import type { CredentialSchema } from "../../projectSchemas/credential"; import type { ConfigBundleSchema } from "../../projectSchemas/config-bundle"; -import type { MemorySchema } from "../../projectSchemas/memory"; +import { MemorySchema } from "../../projectSchemas/memory"; import type { EvaluatorSchema } from "../../projectSchemas/evaluator"; import type { ProjectSpecSchema } from "../../projectSchemas/project"; import z from "zod"; @@ -21,7 +21,7 @@ type CreateProjectInputBase = { skipGit?: boolean; }; -/** Set of flags needed to scaffold a new Runtime-based agent **/ +/** Set of arguments needed to scaffold a new Runtime-based agent. */ export const ScaffoldRuntimeInputSchema = z .object({ runtimeName: AgentNameSchema, @@ -30,7 +30,7 @@ export const ScaffoldRuntimeInputSchema = z framework: z.enum(["strands", "none"]), modelProvider: z.enum(["Bedrock"]), apiKey: z.string().min(1).optional(), - memory: z.enum(["none"]), + memory: MemorySchema.optional(), entrypoint: EntrypointSchema, runtimeVersion: RuntimeVersionSchema.optional(), })