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 45997ff8e..7c458d5da 100644 --- a/src/core/project/manager.test.ts +++ b/src/core/project/manager.test.ts @@ -6,19 +6,19 @@ 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, - type CreateProjectInput, - type DeployResult, - type Project, - type ProjectEvent, +import { resolveRuntimeTemplateShortcut } from "../../handlers/project/shortcuts"; +import type { + CreateProjectInput, + DeployResult, + Project, + ProjectEvent, } from "../../handlers/project/types"; 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[] = []; @@ -101,6 +101,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 bdf4e0503..aa3fc4393 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"]); @@ -83,4 +87,27 @@ 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"); + }); }); diff --git a/src/core/project/templates/fsTree.ts b/src/core/project/templates/fsTree.ts index d388cbe96..7313ed975 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 275d63580..9432a65ac 100644 --- a/src/core/project/templates/runtime.ts +++ b/src/core/project/templates/runtime.ts @@ -61,11 +61,14 @@ const getTemplateResolvers = (assetSource: AssetSource, templateRenderer: Templa if (input.protocol !== undefined && input.protocol !== "HTTP") throw new InputValidationError(`hello-world-python only supports HTTP protocol`); 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)] } }; }, @@ -87,10 +90,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, @@ -105,14 +112,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 8f91a2039..27c0bb693 100644 --- a/src/handlers/project/add/runtime/index.test.ts +++ b/src/handlers/project/add/runtime/index.test.ts @@ -292,6 +292,58 @@ 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); + + const main = await Bun.file(join(projectRoot, "app", "my_agent", "main.py")).text(); + const session = await Bun.file( + join(projectRoot, "app", "my_agent", "memory", "session.py"), + ).text(); + expect(main).toContain("from memory.session import get_memory_session_manager"); + expect(session).toContain('MEMORY_ID = os.getenv("MEMORY_MY_AGENTMEMORY_ID")'); + expect(session.includes("RetrievalConfig")).toBe(expectedStrategies.length > 0); + }); + test.each<[string, string[]]>([ ["missing --name", ["--template", "hello-world-python"]], [ @@ -327,7 +379,7 @@ describe("project add runtime", () => { ], [ "--template and --memory are mutually exclusive", - ["--name", "my_agent", "--template", "hello-world-python", "--memory", "none"], + ["--name", "my_agent", "--template", "strands-python", "--memory", "short"], ], [ "strands-python only supports HTTP", @@ -341,6 +393,7 @@ describe("project add runtime", () => { "hello-world-python only supports HTTP", ["--name", "my_agent", "--template", "hello-world-python", "--protocol", "MCP"], ], + ["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 e386fc9c5..b0b37dcb0 100644 --- a/src/handlers/project/add/runtime/index.ts +++ b/src/handlers/project/add/runtime/index.ts @@ -8,10 +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, - ScaffoldRuntimeInputSchema, -} from "../../types"; + resolveRuntimeTemplateShortcut, +} from "../../shortcuts"; +import { ScaffoldRuntimeInputSchema, type ScaffoldRuntimeInput } from "../../types"; import { RuntimeResourceConfigSchema } from "./types"; export const createAddRuntimeHandler = (config: AddProjectResourceConfig) => @@ -19,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", @@ -48,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", @@ -119,21 +125,22 @@ 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 - ? RUNTIME_TEMPLATE_SHORTCUTS[flags.template!] + const runtimeName = flags.name; + const scaffoldRuntimeInput: ScaffoldRuntimeInput = isTemplate + ? resolveRuntimeTemplateShortcut(flags.template!, runtimeName) : 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 ?? "longAndShortTerm"](runtimeName), entrypoint: "main.py", runtimeVersion: flags.build === "CodeZip" ? "PYTHON_3_14" : undefined, }) - : RUNTIME_TEMPLATE_SHORTCUTS["hello-world-python"]; + : resolveRuntimeTemplateShortcut("hello-world-python", runtimeName); const inputEnvironmentVariables = parseJsonFlag>( "environment-variables", @@ -187,7 +194,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 7252f2048..9498964dd 100644 --- a/src/handlers/project/create/index.ts +++ b/src/handlers/project/create/index.ts @@ -2,11 +2,16 @@ 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, + type ScaffoldRuntimeInput, } from "../types"; import { ProjectNameSchema } from "../../../projectSchemas/project"; import { InputValidationError } from "../../../errors"; @@ -53,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)", @@ -69,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); @@ -85,21 +94,22 @@ 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 - ? RUNTIME_TEMPLATE_SHORTCUTS[flags["template"]!] + const runtimeName = flags["runtime-name"] ?? flags["name"]; + const scaffoldRuntimeInput: ScaffoldRuntimeInput = isTemplate + ? resolveRuntimeTemplateShortcut(flags["template"]!) : 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"] ?? "longAndShortTerm"](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"], @@ -116,7 +126,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 f3e7fc530..4845a8ac3 100644 --- a/src/handlers/project/project.test.ts +++ b/src/handlers/project/project.test.ts @@ -129,6 +129,59 @@ describe("project create", () => { ).rejects.toThrow(/--template and --build are mutually exclusive/); }); + 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); + + const main = await Bun.file(join(projectRoot, "app", "MyAgent", "main.py")).text(); + const session = await Bun.file( + join(projectRoot, "app", "MyAgent", "memory", "session.py"), + ).text(); + expect(main).toContain("from memory.session import get_memory_session_manager"); + expect(session).toContain('MEMORY_ID = os.getenv("MEMORY_MYAGENTMEMORY_ID")'); + expect(session.includes("RetrievalConfig")).toBe(expectedStrategies.length > 0); + }); + test("scaffolds from explicit custom flags", async () => { const directory = await inTempDirectory(); await run([ @@ -162,32 +215,41 @@ 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([ + ["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); - }); + expect(existsSync(join(directory, "MyProject"))).toBe(false); + }, + ); test("rejects an API key with the Bedrock model provider before scaffolding", async () => { const directory = await inTempDirectory(); diff --git a/src/handlers/project/shortcuts.ts b/src/handlers/project/shortcuts.ts new file mode 100644 index 000000000..cb065cbe1 --- /dev/null +++ b/src/handlers/project/shortcuts.ts @@ -0,0 +1,89 @@ +import { + DEFAULT_EPISODIC_REFLECTION_NAMESPACE_TEMPLATES, + DEFAULT_STRATEGY_NAMESPACE_TEMPLATES, + type Memory, +} from "../../projectSchemas/memory"; +import 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", + build: "CodeZip", + language: "Python", + framework: "none", + modelProvider: "Bedrock", + entrypoint: "main.py", + runtimeVersion: "PYTHON_3_14", + }, + "hello-world-python-container": { + runtimeName: "hello_world", + build: "Container", + language: "Python", + framework: "none", + modelProvider: "Bedrock", + entrypoint: "main.py", + }, + "strands-python": { + runtimeName: "strands_agent", + build: "CodeZip", + language: "Python", + framework: "strands", + modelProvider: "Bedrock", + entrypoint: "main.py", + runtimeVersion: "PYTHON_3_14", + memory: "longAndShortTerm", + }, +} as const satisfies Record; + +export type RuntimeTemplateShortcutName = keyof typeof RUNTIME_TEMPLATE_SHORTCUTS; + +export const RUNTIME_TEMPLATE_SHORTCUT_NAMES = Object.keys( + RUNTIME_TEMPLATE_SHORTCUTS, +) as unknown as readonly [RuntimeTemplateShortcutName, ...RuntimeTemplateShortcutName[]]; + +export function resolveRuntimeTemplateShortcut( + name: RuntimeTemplateShortcutName, + runtimeName: string = RUNTIME_TEMPLATE_SHORTCUTS[name].runtimeName, +): ScaffoldRuntimeInput { + const selected: RuntimeTemplateShortcut = RUNTIME_TEMPLATE_SHORTCUTS[name]; + const { memory: memoryShortcut, ...shortcut } = selected; + const memory = memoryShortcut ? MEMORY_SHORTCUTS[memoryShortcut](runtimeName) : undefined; + return { + ...shortcut, + runtimeName, + ...(memory && { memory }), + }; +} diff --git a/src/handlers/project/types.ts b/src/handlers/project/types.ts index c96b2e7af..6d6dddb09 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"; @@ -12,44 +12,6 @@ import { RuntimeVersionSchema } from "../../projectSchemas/constants"; import type { AgentCoreGateway, AgentCoreGatewayTarget } from "../../projectSchemas/gateway"; import type { PolicyEngineSchema, PolicySchema } from "../../projectSchemas/policy"; -export const RUNTIME_TEMPLATE_SHORTCUTS = { - "hello-world-python": { - runtimeName: "hello_world", - build: "CodeZip", - language: "Python", - framework: "none", - modelProvider: "Bedrock", - memory: "none", - entrypoint: "main.py", - runtimeVersion: "PYTHON_3_14", - }, - "hello-world-python-container": { - runtimeName: "hello_world", - build: "Container", - language: "Python", - framework: "none", - modelProvider: "Bedrock", - memory: "none", - entrypoint: "main.py", - }, - "strands-python": { - runtimeName: "strands_agent", - build: "CodeZip", - language: "Python", - framework: "strands", - modelProvider: "Bedrock", - memory: "none", - entrypoint: "main.py", - runtimeVersion: "PYTHON_3_14", - }, -} as const satisfies Record; - -export type RuntimeTemplateShortcutName = keyof typeof RUNTIME_TEMPLATE_SHORTCUTS; - -export const RUNTIME_TEMPLATE_SHORTCUT_NAMES = Object.keys( - RUNTIME_TEMPLATE_SHORTCUTS, -) as unknown as readonly [RuntimeTemplateShortcutName, ...RuntimeTemplateShortcutName[]]; - type CreateProjectInputBase = { /** The name of the project; also the directory it is scaffolded into. */ name: string; @@ -59,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, @@ -68,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(), })