Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 50 additions & 5 deletions src/cli/primitives/PolicyPrimitive.ts
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import { ResourceNotFoundError, ValidationError, findConfigRoot, serializeResult, toError } from '../../lib';
import type { Result } from '../../lib/result';
import type { Policy } from '../../schema';
import type { AgentCoreGatewayTarget, AgentCoreProjectSpec, Policy } from '../../schema';
import { EnforcementModeSchema, PolicySchema, ValidationModeSchema } from '../../schema';
import { detectRegion } from '../aws';
import { getPolicyGeneration, startPolicyGeneration } from '../aws/policy-generation';
Expand All @@ -24,6 +24,38 @@ import type { AddResult, AddScreenComponent, RemovableResource } from './types';
import type { Command } from '@commander-js/extra-typings';
import { existsSync, readFileSync } from 'fs';

/** Return a tool name only when the target exposes exactly one known tool. */
function singleKnownToolName(target: AgentCoreGatewayTarget | undefined): string | undefined {
if (!target) return undefined;

const toolNames = [
...(target.toolDefinitions ?? []).map(tool => tool.name),
...(target.configurations ?? []).map(configuration => configuration.name),
];
return toolNames.length === 1 ? toolNames[0] : undefined;
}

export function resolvePolicyTarget(
project: AgentCoreProjectSpec,
targetName: string,
gatewayName?: string
): { gatewayName?: string; toolName?: string } {
const matchingGateways = project.agentCoreGateways.filter(gateway =>
gateway.targets.some(target => target.name === targetName)
);
if (!gatewayName && matchingGateways.length > 1) {
throw new ValidationError(
`Target "${targetName}" exists on multiple gateways: ${matchingGateways.map(g => g.name).join(', ')}. Use --gateway <name> to specify one.`
);
}

const resolvedGatewayName = gatewayName ?? matchingGateways[0]?.name;
const target = project.agentCoreGateways
.find(gateway => gateway.name === resolvedGatewayName)
?.targets.find(candidate => candidate.name === targetName);
return { gatewayName: resolvedGatewayName, toolName: singleKnownToolName(target) };
}

export interface AddPolicyOptions {
name: string;
engine: string;
Expand Down Expand Up @@ -417,17 +449,30 @@ export class PolicyPrimitive extends BasePrimitive<AddPolicyOptions, RemovablePo

let resolvedGatewayArn: string | undefined;
let resolvedTargetName: string | undefined = cliOptions.target;
if (cliOptions.gateway) {
let resolvedToolName: string | undefined;
let resolvedGatewayName: string | undefined = cliOptions.gateway;

// A target name identifies its gateway in the project spec. Resolve that
// relationship when --gateway is omitted so target-scoped policies can
// still be constrained to the deployed gateway resource.
if (cliOptions.target) {
const project = await this.readProjectSpec();
const resolved = resolvePolicyTarget(project, cliOptions.target, resolvedGatewayName);
resolvedGatewayName = resolved.gatewayName;
resolvedToolName = resolved.toolName;
}

if (resolvedGatewayName) {
try {
const deployedState = await this.configIO.readDeployedState();
const gateway = findDeployedGateway(deployedState, cliOptions.gateway);
const gateway = findDeployedGateway(deployedState, resolvedGatewayName);
if (gateway?.gatewayArn) {
resolvedGatewayArn = gateway.gatewayArn;
if (!resolvedTargetName) {
const targetNames = Object.keys(gateway.targets ?? {});
if (targetNames.length > 1) {
throw new ValidationError(
`Multiple targets found on gateway "${cliOptions.gateway}": ${targetNames.join(', ')}. Use --target <name> to specify one.`
`Multiple targets found on gateway "${resolvedGatewayName}": ${targetNames.join(', ')}. Use --target <name> to specify one.`
);
}
resolvedTargetName = targetNames[0];
Expand All @@ -448,7 +493,7 @@ export class PolicyPrimitive extends BasePrimitive<AddPolicyOptions, RemovablePo
effect: policyEffect,
dataPath: cliOptions.formDataPath ?? defaultDataPathForEffect(policyEffect),
},
{ targetName: resolvedTargetName, gatewayArn: resolvedGatewayArn }
{ targetName: resolvedTargetName, toolName: resolvedToolName, gatewayArn: resolvedGatewayArn }
);
// Output-phase effects (suppressOutput) must register on RETURN_OUTPUT.
effectiveOptions = {
Expand Down
55 changes: 54 additions & 1 deletion src/cli/primitives/__tests__/PolicyPrimitive.test.ts
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import type { AgentCoreProjectSpec, Policy, PolicyEngine } from '../../../schema';
import { PolicyPrimitive } from '../PolicyPrimitive';
import { PolicyPrimitive, resolvePolicyTarget } from '../PolicyPrimitive';
import { beforeEach, describe, expect, it, vi } from 'vitest';

const engine: PolicyEngine = { name: 'eng', policies: [] };
Expand Down Expand Up @@ -128,6 +128,59 @@ describe('PolicyPrimitive — enforcementMode', () => {
});
});

describe('resolvePolicyTarget', () => {
it('resolves a unique connector target and its tool name', () => {
const project = {
...defaultProject,
agentCoreGateways: [
{
name: 'toolgw',
protocolType: 'MCP' as const,
authorizerType: 'NONE' as const,
enableSemanticSearch: true,
exceptionLevel: 'NONE' as const,
targets: [
{
name: 'websearch',
targetType: 'connector' as const,
connectorId: 'web-search' as const,
configurations: [{ name: 'WebSearch' }],
},
],
},
],
};

expect(resolvePolicyTarget(project, 'websearch')).toEqual({ gatewayName: 'toolgw', toolName: 'WebSearch' });
});

it('requires an explicit gateway when a target name is ambiguous', () => {
const project = {
...defaultProject,
agentCoreGateways: [
{
name: 'first',
protocolType: 'MCP' as const,
authorizerType: 'NONE' as const,
enableSemanticSearch: true,
exceptionLevel: 'NONE' as const,
targets: [{ name: 'shared', targetType: 'connector' as const }],
},
{
name: 'second',
protocolType: 'MCP' as const,
authorizerType: 'NONE' as const,
enableSemanticSearch: true,
exceptionLevel: 'NONE' as const,
targets: [{ name: 'shared', targetType: 'connector' as const }],
},
],
};

expect(() => resolvePolicyTarget(project, 'shared')).toThrow('Use --gateway <name> to specify one.');
});
});

describe('PolicyPrimitive — --generate gateway resolution', () => {
let primitive: PolicyPrimitive;

Expand Down
6 changes: 6 additions & 0 deletions src/cli/tui/screens/policy/__tests__/synthesize-cedar.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,12 @@ describe('synthesizeCedar', () => {
expect(result).toContain('resource == AgentCore::Gateway::"arn:aws:agentcore:us-east-1:123456:gateway/gw-abc"');
});

it('uses the target tool name when provided', () => {
const result = synthesizeCedar(baseForm, { targetName: 'websearch', toolName: 'WebSearch' });
expect(result).toContain('action == AgentCore::Action::"websearch___WebSearch"');
expect(result).not.toContain('POST:/invocations');
});

it('uses custom dataPath', () => {
const form: GuardrailFormConfig = { ...baseForm, dataPath: 'context.output.response' };
const result = synthesizeCedar(form);
Expand Down
8 changes: 6 additions & 2 deletions src/cli/tui/screens/policy/synthesize-cedar.ts
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,8 @@ const DEFAULT_THRESHOLDS: Record<GuardrailCategoryType, number> = {
*/
export interface SynthesizeCedarOptions {
targetName?: string;
/** Tool exposed by the target. Required for a tool-scoped action reference. */
toolName?: string;
gatewayArn?: string;
}

Expand All @@ -34,10 +36,12 @@ export function synthesizeCedar(form: GuardrailFormConfig, options: SynthesizeCe
return '// No guardrail rules configured';
}

const { targetName, gatewayArn } = options;
const { targetName, toolName, gatewayArn } = options;
const fn = GUARDRAIL_FUNCTION_MAP[form.category];
const gwRef = gatewayArn ? `resource == AgentCore::Gateway::"${gatewayArn}"` : 'resource is AgentCore::Gateway';
const actionRef = targetName ? `action == AgentCore::Action::"${targetName}___POST:/invocations"` : 'action';
const actionRef = targetName
? `action == AgentCore::Action::"${targetName}___${toolName ?? 'POST:/invocations'}"`
: 'action';
const dataPath = form.dataPath || defaultDataPathForEffect(form.effect);
const threshold = DEFAULT_THRESHOLDS[form.category];
// permit allows below threshold; forbid and suppressOutput block above it.
Expand Down
Loading