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
203 changes: 60 additions & 143 deletions app/api/claude/route.ts
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
import Anthropic from "@anthropic-ai/sdk";
import OpenAI from "openai";
// ABOUTME: API route for LLM-powered music generation
// ABOUTME: Uses Vercel AI SDK to abstract over Anthropic and OpenAI providers

import { streamText } from "ai";
import { NextRequest } from "next/server";
import { buildSystemPrompt } from "@/lib/systemPrompt";
import {
Expand All @@ -8,6 +10,12 @@ import {
buildRetryPrompt,
} from "@/lib/creativeDirectives";
import { validateStrudelCode } from "@/lib/validateOutput";
import {
getProviderFromModel,
validateApiKeyForProvider,
createModel,
type Provider,
} from "@/lib/ai/providers";
import type { GenerationMode, ChatMessage } from "@/lib/sessionStore";

// extracts code from response for validation
Expand Down Expand Up @@ -80,98 +88,40 @@ function formatChatHistory(chat: ChatMessage[]): string {
return `\nprevious conversation:\n${formatted}\n`;
}

// streams claude response
async function streamClaude(
client: Anthropic,
// stream text and return accumulated result, outputting in the expected SSE format
async function streamToClient(
provider: Provider,
modelId: string,
apiKey: string,
systemPrompt: string,
userPrompt: string,
maxTokens: number,
controller: ReadableStreamDefaultController,
encoder: TextEncoder,
model: string
encoder: TextEncoder
): Promise<string> {
const stream = await client.messages.stream({
model: model,
max_tokens: maxTokens,
system: systemPrompt,
messages: [{ role: "user", content: userPrompt }],
});

let fullText = "";
const model = createModel(provider, modelId, apiKey);

for await (const event of stream) {
if (event.type === "content_block_delta") {
const delta = event.delta as { type: string; text?: string };
if (delta.type === "text_delta" && delta.text) {
fullText += delta.text;
const data = JSON.stringify({
type: "content_block_delta",
delta: { text: delta.text },
});
controller.enqueue(encoder.encode(`data: ${data}\n\n`));
}
}
}

return fullText;
}

// streams OpenAI response
async function streamOpenAI(
client: OpenAI,
systemPrompt: string,
userPrompt: string,
maxTokens: number,
controller: ReadableStreamDefaultController,
encoder: TextEncoder,
model: string
): Promise<string> {
const stream = await client.chat.completions.create({
model: model,
max_completion_tokens: maxTokens,
messages: [
{ role: "system", content: systemPrompt },
{ role: "user", content: userPrompt },
],
stream: true,
const result = streamText({
model,
system: systemPrompt,
prompt: userPrompt,
maxOutputTokens: maxTokens,
});

let fullText = "";

for await (const chunk of stream) {
const delta = chunk.choices[0]?.delta?.content;
if (delta) {
fullText += delta;
const data = JSON.stringify({
type: "content_block_delta",
delta: { text: delta },
});
controller.enqueue(encoder.encode(`data: ${data}\n\n`));
}
for await (const chunk of result.textStream) {
fullText += chunk;
const data = JSON.stringify({
type: "content_block_delta",
delta: { text: chunk },
});
controller.enqueue(encoder.encode(`data: ${data}\n\n`));
}

return fullText;
}

// Helper to determine provider from model ID
function getProviderFromModel(model: string): "anthropic" | "openai" {
if (model.startsWith("gpt-") || model.startsWith("o1") || model.startsWith("o3")) {
return "openai";
}
return "anthropic";
}

// Helper to detect API key type
function detectApiKeyProvider(apiKey: string): "anthropic" | "openai" | "unknown" {
if (apiKey.startsWith("sk-ant-")) {
return "anthropic";
}
if (apiKey.startsWith("sk-") || apiKey.startsWith("sk-proj-")) {
return "openai";
}
return "unknown";
}

export async function POST(req: NextRequest) {
try {
const body = await req.json();
Expand Down Expand Up @@ -202,18 +152,14 @@ export async function POST(req: NextRequest) {
const apiKey = requestApiKey;

// Use client-provided model if available, otherwise default
const model = requestModel || "claude-sonnet-4-20250514";
const provider = getProviderFromModel(model);

// Detect API key type and validate it matches the selected provider
const keyProvider = detectApiKeyProvider(apiKey);
if (keyProvider !== "unknown" && keyProvider !== provider) {
const providerName = provider === "openai" ? "OpenAI" : "Anthropic";
const keyProviderName = keyProvider === "openai" ? "OpenAI" : "Anthropic";
const modelId = requestModel || "claude-sonnet-4-20250514";
const provider = getProviderFromModel(modelId);

// Validate API key matches the selected provider
const keyError = validateApiKeyForProvider(apiKey, provider);
if (keyError) {
return new Response(
JSON.stringify({
error: `You selected an ${providerName} model but provided an ${keyProviderName} API key. Please update your API key in Settings to match the selected model.`,
}),
JSON.stringify({ error: keyError }),
{
status: 400,
headers: { "Content-Type": "application/json" },
Expand All @@ -228,10 +174,6 @@ export async function POST(req: NextRequest) {
});
}

// Create appropriate client based on provider
const anthropicClient = provider === "anthropic" ? new Anthropic({ apiKey }) : null;
const openaiClient = provider === "openai" ? new OpenAI({ apiKey }) : null;

// get balanced config values
const config = getConfigValues();

Expand All @@ -242,7 +184,7 @@ export async function POST(req: NextRequest) {
let userPrompt: string;

if (mode === "edit" && currentCode && currentCode.trim()) {
// edit mode: include current code context
// edit mode: include current code context and chat history
const chatContext = formatChatHistory(truncateChatHistory(chatHistory, 6));
userPrompt = buildEditContext(currentCode, prompt) + chatContext;
} else {
Expand All @@ -259,30 +201,16 @@ export async function POST(req: NextRequest) {
async start(controller) {
try {
// first attempt - streaming
let fullText: string;
if (provider === "openai" && openaiClient) {
fullText = await streamOpenAI(
openaiClient,
systemPrompt,
userPrompt,
maxTokens,
controller,
encoder,
model
);
} else if (anthropicClient) {
fullText = await streamClaude(
anthropicClient,
systemPrompt,
userPrompt,
maxTokens,
controller,
encoder,
model
);
} else {
throw new Error("No valid API client configured");
}
const fullText = await streamToClient(
provider,
modelId,
apiKey,
systemPrompt,
userPrompt,
maxTokens,
controller,
encoder
);

// validate the output
const extractedCode = extractCodeForValidation(fullText);
Expand All @@ -305,7 +233,7 @@ export async function POST(req: NextRequest) {

// build retry prompt with issues
const retryPrompt = buildRetryPrompt(prompt, validation.issues);

// for edit mode, include the current code in retry
let retryUserPrompt: string;
if (mode === "edit" && currentCode && currentCode.trim()) {
Expand All @@ -319,35 +247,24 @@ export async function POST(req: NextRequest) {
controller.enqueue(encoder.encode(`data: ${clearMsg}\n\n`));

// second attempt - also streaming
if (provider === "openai" && openaiClient) {
await streamOpenAI(
openaiClient,
systemPrompt,
retryUserPrompt,
maxTokens,
controller,
encoder,
model
);
} else if (anthropicClient) {
await streamClaude(
anthropicClient,
systemPrompt,
retryUserPrompt,
maxTokens,
controller,
encoder,
model
);
}
await streamToClient(
provider,
modelId,
apiKey,
systemPrompt,
retryUserPrompt,
maxTokens,
controller,
encoder
);
}
}

controller.enqueue(encoder.encode("data: [DONE]\n\n"));
controller.close();
} catch (error) {
let errorMessage = "stream error";

if (error instanceof Error) {
console.error("API Error:", error.message, error);
// Check for authentication/invalid API key errors
Expand All @@ -368,7 +285,7 @@ export async function POST(req: NextRequest) {
} else {
console.error("Unknown API Error:", error);
}

const data = JSON.stringify({
type: "error",
error: { message: errorMessage },
Expand Down
Loading