Skip to content
Draft
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
28 changes: 11 additions & 17 deletions src/app/api/projects/[projectId]/runs/route.ts
Original file line number Diff line number Diff line change
Expand Up @@ -109,16 +109,20 @@ function normalizeSchedulerDrainRequest(value: unknown) {
value && typeof value === "object" && !Array.isArray(value)
? (value as Record<string, unknown>)
: {};
const researchRunId = readOptionalString(body.researchRunId);
const executionProfileId = readOptionalString(body.executionProfileId);
const threadId = readOptionalString(body.threadId);
if (!researchRunId || !executionProfileId || !threadId) {
throw new Error(
"Background run drain requires researchRunId, executionProfileId, and threadId.",
);
}
return {
...(readOptionalString(body.threadId) ? { threadId: readOptionalString(body.threadId) } : {}),
researchRunId,
executionProfileId,
threadId,
...(readOptionalString(body.owner) ? { owner: readOptionalString(body.owner) } : {}),
...(readOptionalString(body.goal) ? { goal: readOptionalString(body.goal) } : {}),
...(readOptionalStringArray(body.approvedToolIds)
? { approvedToolIds: readOptionalStringArray(body.approvedToolIds) }
: {}),
...(readOptionalStringArray(body.commandAllowlist)
? { commandAllowlist: readOptionalStringArray(body.commandAllowlist) }
: {}),
...(readOptionalPositiveInteger(body.concurrency)
? { concurrency: readOptionalPositiveInteger(body.concurrency) }
: {}),
Expand All @@ -135,16 +139,6 @@ function readOptionalString(value: unknown) {
return typeof value === "string" && value.trim() ? value.trim() : undefined;
}

function readOptionalStringArray(value: unknown) {
if (!Array.isArray(value)) {
return undefined;
}
const items = value
.filter((item): item is string => typeof item === "string" && item.trim().length > 0)
.map((item) => item.trim());
return items.length > 0 ? items : undefined;
}

function readOptionalPositiveInteger(value: unknown) {
if (value === undefined) {
return undefined;
Expand Down
68 changes: 48 additions & 20 deletions src/mastra/agent-controller/agent-controller.ts
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,8 @@ import {
type ToolCategory,
} from "@mastra/core/agent-controller";
import { MastraModelGateway, type MastraModelGatewayInterface } from "@mastra/core/llm";
import { parse } from "llm-strings";
import { z } from "zod";

import { asModelConnectionString, modelProviderFromHost } from "../../lib/models";
import { compactResearchStateSchema, emptyCompactResearchState } from "../../lib/research-terminal";
import { classifySecurityCapability, LOCAL_TOOL_IDS } from "../../lib/tools/catalog";
import {
Expand All @@ -21,7 +19,11 @@ import {
} from "../agents/security-research";
import { securityResearchAgent } from "../agents/security-research-agent";
import { resolveSecurityResearchControllerMemory } from "../config/memory";
import { resolveSecurityResearchMastraModelUri } from "../config/model";
import {
getMastraChatModel,
getMastraModelRuntimeOptions,
resolveSecurityResearchMastraModelUri,
} from "../config/model";
import { mastraStorage } from "../config/storage";
import { securityResearchWorkspace } from "../config/workspace";
import { selectSecurityResearchTools } from "../tools";
Expand All @@ -43,6 +45,14 @@ class SecurityResearchNativeModelGateway extends MastraModelGateway {
gateway: this.id,
npm: "@ai-sdk/anthropic",
},
profile: {
apiKeyEnvVar: "MODEL_DEFAULT",
name: "Pinned ExploitHunter profile",
models: [],
docUrl: "https://ExploitHunter.app/",
gateway: this.id,
npm: "ai",
},
};
}

Expand All @@ -51,11 +61,7 @@ class SecurityResearchNativeModelGateway extends MastraModelGateway {
}

async getApiKey() {
const apiKey = process.env.ANTHROPIC_API_KEY?.trim();
if (!apiKey) {
throw new Error("ANTHROPIC_API_KEY is required for native Anthropic AgentController models.");
}
return apiKey;
return process.env.ANTHROPIC_API_KEY?.trim() || "pinned-profile";
}

resolveLanguageModel({
Expand All @@ -69,13 +75,44 @@ class SecurityResearchNativeModelGateway extends MastraModelGateway {
apiKey: string;
headers?: Record<string, string>;
}) {
if (providerId === "profile") {
const modelUri = Buffer.from(modelId, "base64url").toString("utf8");
return applyPinnedControllerModelOptions(
modelUri,
getMastraChatModel(modelUri) as Record<string, unknown>,
) as never;
}
if (providerId !== "anthropic") {
throw new Error(`Unsupported native AgentController provider: ${providerId}`);
}
return createAnthropic({ apiKey, headers }).chat(modelId);
}
}

export function applyPinnedControllerModelOptions(
modelUri: string,
model: Record<string, unknown>,
) {
const runtime = getMastraModelRuntimeOptions(modelUri);
return new Proxy(model, {
get(target, property, receiver) {
const value = Reflect.get(target, property, receiver);
if ((property === "doGenerate" || property === "doStream") && typeof value === "function") {
return (options: Record<string, unknown>) =>
value.call(target, {
...options,
...runtime.modelSettings,
providerOptions: {
...((options.providerOptions as Record<string, unknown> | undefined) ?? {}),
...runtime.providerOptions,
},
});
}
return typeof value === "function" ? value.bind(target) : value;
},
});
}

export const securityResearchControllerGateways: MastraModelGatewayInterface[] = [
new SecurityResearchNativeModelGateway(),
];
Expand Down Expand Up @@ -119,16 +156,7 @@ const DEFAULT_SECURITY_RESEARCH_MODEL_URI = resolveSecurityResearchMastraModelUr
);

export function toMastraGatewayModelId(modelUri: string) {
const parsed = parse(asModelConnectionString(modelUri));
const provider =
modelProviderFromHost(parsed.hostAlias) ??
modelProviderFromHost(parsed.host) ??
parsed.hostAlias ??
parsed.host;
if (provider === "anthropic" && parsed.model === "claude-opus-5") {
return `${SECURITY_RESEARCH_NATIVE_GATEWAY_ID}/${provider}/${parsed.model}`;
}
return provider ? `${provider}/${parsed.model}` : parsed.model;
return `${SECURITY_RESEARCH_NATIVE_GATEWAY_ID}/profile/${Buffer.from(modelUri).toString("base64url")}`;
}

const DEFAULT_SECURITY_RESEARCH_MODEL = toMastraGatewayModelId(DEFAULT_SECURITY_RESEARCH_MODEL_URI);
Expand Down Expand Up @@ -206,8 +234,8 @@ export const securityResearchControllerSubagents: AgentControllerSubagent[] =
id: stage,
name: securityResearchControllerSubagentStageNames[stage],
description: securityResearchControllerSubagentDescriptions[stage],
instructions: async () =>
String(await securityResearchStageAgents[agentKey].getInstructions()),
instructions: async ({ requestContext }) =>
String(await securityResearchStageAgents[agentKey].getInstructions({ requestContext })),
allowedControllerTools: [...getSecurityResearchStageToolIds(stage)],
defaultModelId: DEFAULT_SECURITY_RESEARCH_MODEL,
maxSteps: stage === "hunt" || stage === "trace" ? 24 : 12,
Expand Down
32 changes: 28 additions & 4 deletions src/mastra/agents/security-research/stage-agents.ts
Original file line number Diff line number Diff line change
Expand Up @@ -288,11 +288,16 @@ function stageModel(definition: StageDefinition) {
async function stageInstructions(
definition: StageDefinition,
capabilities: readonly string[] = [],
pinnedSkills: readonly { id: string; revision: string; detail?: string }[] = [],
pinnedProfile = false,
) {
const skillInstructions = await formatSecurityResearchSkillInstructions(
definition.skillHints,
{ capabilities },
);
const skillInstructions = pinnedProfile
? pinnedSkills
.filter((skill) => definition.skillHints.includes(skill.id))
.filter((skill) => skill.detail)
.map((skill) => `## ${skill.id} [${skill.revision}]\n${skill.detail}`)
.join("\n\n")
: await formatSecurityResearchSkillInstructions(definition.skillHints, { capabilities });
return `${commonStageInstructions(definition.stage)}

Stage role:
Expand Down Expand Up @@ -320,6 +325,8 @@ function createStageAgent(definition: StageDefinition) {
stageInstructions(
definition,
readRuntimeSkillCapabilities(requestContext.get("runtimeSkillCapabilities")),
readPinnedSkillRefs(requestContext.get("selectedSkillRefs")),
typeof requestContext.get("researchExecutionProfileId") === "string",
),
model: stageModel(definition),
memory: createSecurityResearchStageMemory(
Expand All @@ -345,6 +352,23 @@ function createStageAgent(definition: StageDefinition) {
});
}

function readPinnedSkillRefs(value: unknown) {
if (!Array.isArray(value)) return [];
return value.flatMap((item) => {
if (!item || typeof item !== "object" || Array.isArray(item)) return [];
const skill = item as Record<string, unknown>;
return typeof skill.id === "string" && typeof skill.revision === "string"
? [
{
id: skill.id,
revision: skill.revision,
...(typeof skill.detail === "string" ? { detail: skill.detail } : {}),
},
]
: [];
});
}

function readRuntimeSkillCapabilities(value: unknown) {
if (!Array.isArray(value)) return [];
return value.filter((item): item is string => typeof item === "string");
Expand Down
Loading