Skip to content
Closed
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
170 changes: 170 additions & 0 deletions src/providers/claudeCodeSdkProvider.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,170 @@
import * as vscode from 'vscode';
import * as path from 'path';
import * as fs from 'fs';
import * as os from 'os';
import { query } from '@anthropic-ai/claude-code';
import { LLMProvider, LLMProviderOptions, LLMMessage, LLMStreamCallback } from './llmProvider';
import { Logger } from '../services/logger';

export class ClaudeCodeSdkProvider extends LLMProvider {
private workingDirectory: string = '';
private currentSessionId: string | null = null;
private modelId: string = 'claude-sonnet-4-20250514';

constructor(outputChannel: vscode.OutputChannel) {
super(outputChannel);
this.initializationPromise = this.initialize();
}

async initialize(): Promise<void> {
if (this.isInitialized) {
return;
}
try {
Logger.info('Starting Claude Code SDK provider initialization...');
await this.setupWorkingDirectory();
await this.loadConfiguration();
this.isInitialized = true;
Logger.info('Claude Code SDK provider initialized successfully');
} catch (error) {
Logger.error(`Failed to initialize Claude Code SDK provider: ${error}`);
this.initializationPromise = null;
this.isInitialized = false;
throw error;
}
}

private async setupWorkingDirectory(): Promise<void> {
try {
const workspaceRoot = vscode.workspace.workspaceFolders?.[0]?.uri.fsPath;
if (workspaceRoot) {
const superdesignDir = path.join(workspaceRoot, '.superdesign');
if (!fs.existsSync(superdesignDir)) {
fs.mkdirSync(superdesignDir, { recursive: true });
}
this.workingDirectory = superdesignDir;
} else {
const tempDir = path.join(os.tmpdir(), 'superdesign-claude-code-sdk');
if (!fs.existsSync(tempDir)) {
fs.mkdirSync(tempDir, { recursive: true });
}
this.workingDirectory = tempDir;
}
} catch (error) {
Logger.error(`Failed to setup working directory: ${error}`);
this.workingDirectory = process.cwd();
}
}

private async loadConfiguration(): Promise<void> {
const config = vscode.workspace.getConfiguration('superdesign');
const modelId = config.get<string>('claudeCodeModelId');
if (modelId) {
this.modelId = modelId;
}
}

async query(
prompt: string,
options?: Partial<LLMProviderOptions>,
abortController?: AbortController,
onMessage?: LLMStreamCallback
): Promise<LLMMessage[]> {
Logger.info('Starting Claude Code SDK query');
await this.ensureInitialized();

const messages: LLMMessage[] = [];

try {
const sdkOptions = {
cwd: options?.cwd ?? this.workingDirectory,
maxTurns: options?.maxTurns ?? 10,
allowedTools: options?.allowedTools ?? [
'Read', 'Write', 'Edit', 'MultiEdit', 'Bash', 'LS', 'Grep', 'Glob',
],
permissionMode: (options?.permissionMode ?? 'acceptEdits') as 'acceptEdits' | 'default' | 'bypassPermissions' | 'plan',
customSystemPrompt: options?.customSystemPrompt,
model: this.modelId,
resume: this.currentSessionId ?? options?.resume,
};

const stream = query({ prompt, abortController, options: sdkOptions });

for await (const sdkMessage of stream) {
const message: LLMMessage = sdkMessage as unknown as LLMMessage;
messages.push(message);

if (sdkMessage.session_id && !this.currentSessionId) {
this.currentSessionId = sdkMessage.session_id;
}

if (onMessage) {
try {
onMessage(message);
} catch (callbackError) {
Logger.error(`Streaming callback error: ${callbackError}`);
}
}
}

Logger.info(`Claude Code SDK query completed with ${messages.length} messages`);
return messages;
} catch (error) {
Logger.error(`Claude Code SDK query failed: ${error}`);
throw error;
}
}

isReady(): boolean {
return this.isInitialized;
}

async waitForInitialization(): Promise<boolean> {
try {
await this.ensureInitialized();
return true;
} catch (error) {
Logger.error(`Claude Code SDK provider initialization failed: ${error}`);
return false;
}
}

getWorkingDirectory(): string {
return this.workingDirectory;
}

hasValidConfiguration(): boolean {
return true;
}

async refreshConfiguration(): Promise<boolean> {
try {
await this.loadConfiguration();
return true;
} catch (error) {
Logger.error(`Failed to refresh Claude Code SDK configuration: ${error}`);
return false;
}
}

isAuthError(errorMessage: string): boolean {
const authErrorPatterns = [
'authentication failed',
'unauthorized',
'access denied',
'permission denied',
'invalid api key',
'api key',
];
const lower = errorMessage.toLowerCase();
return authErrorPatterns.some(p => lower.includes(p));
}

getProviderName(): string {
return 'Claude Code SDK';
}

getProviderType(): 'api' | 'binary' | 'sdk' {
return 'sdk';
}
}
7 changes: 4 additions & 3 deletions src/providers/llmProvider.ts
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,7 @@ export abstract class LLMProvider {
abstract refreshConfiguration(): Promise<boolean>;
abstract isAuthError(errorMessage: string): boolean;
abstract getProviderName(): string;
abstract getProviderType(): 'api' | 'binary';
abstract getProviderType(): 'api' | 'binary' | 'sdk';

protected async ensureInitialized(): Promise<void> {
if (this.initializationPromise) {
Expand All @@ -70,7 +70,8 @@ export abstract class LLMProvider {

export enum LLMProviderType {
CLAUDE_API = 'claude-api',
CLAUDE_CODE = 'claude-code'
CLAUDE_CODE = 'claude-code',
CLAUDE_CODE_SDK = 'claude-code-sdk',
}

export interface LLMProviderConfig {
Expand All @@ -79,4 +80,4 @@ export interface LLMProviderConfig {
claudeCodePath?: string;
modelId?: string;
thinkingBudgetTokens?: number;
}
}
54 changes: 23 additions & 31 deletions src/providers/llmProviderFactory.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import * as vscode from 'vscode';
import { LLMProvider, LLMProviderType } from './llmProvider';
import { ClaudeApiProvider } from './claudeApiProvider';
import { ClaudeCodeProvider } from './claudeCodeProvider';
import { ClaudeCodeSdkProvider } from './claudeCodeSdkProvider';
import { Logger } from '../services/logger';

export class LLMProviderFactory {
Expand All @@ -19,12 +20,10 @@ export class LLMProviderFactory {
}

async getProvider(providerType?: LLMProviderType): Promise<LLMProvider> {
// If no provider type specified, get from configuration
if (!providerType) {
providerType = this.getConfiguredProviderType();
}

// Check if provider is already initialized
if (this.providers.has(providerType)) {
const provider = this.providers.get(providerType)!;
if (provider.isReady()) {
Expand All @@ -33,29 +32,26 @@ export class LLMProviderFactory {
}
}

// Create new provider
const provider = await this.createProvider(providerType);

// Initialize the provider
await provider.waitForInitialization();

// Store the provider
this.providers.set(providerType, provider);
this.currentProvider = provider;

return provider;
}

private async createProvider(providerType: LLMProviderType): Promise<LLMProvider> {
Logger.info(`Creating provider of type: ${providerType}`);

switch (providerType) {
case LLMProviderType.CLAUDE_API:
return new ClaudeApiProvider(this.outputChannel);

case LLMProviderType.CLAUDE_CODE:
return new ClaudeCodeProvider(this.outputChannel);


case LLMProviderType.CLAUDE_CODE_SDK:
return new ClaudeCodeSdkProvider(this.outputChannel);

default:
throw new Error(`Unknown provider type: ${providerType}`);
}
Expand All @@ -64,9 +60,10 @@ export class LLMProviderFactory {
private getConfiguredProviderType(): LLMProviderType {
const config = vscode.workspace.getConfiguration('superdesign');
const providerType = config.get<string>('llmProvider', 'claude-api');

// Map string to enum

switch (providerType.toLowerCase()) {
case 'claude-code-sdk':
return LLMProviderType.CLAUDE_CODE_SDK;
case 'claude-code':
return LLMProviderType.CLAUDE_CODE;
case 'claude-api':
Expand All @@ -83,7 +80,6 @@ export class LLMProviderFactory {
if (!this.currentProvider) {
return false;
}

try {
return await this.currentProvider.refreshConfiguration();
} catch (error) {
Expand All @@ -94,12 +90,8 @@ export class LLMProviderFactory {

async switchProvider(providerType: LLMProviderType): Promise<LLMProvider> {
Logger.info(`Switching to provider: ${providerType}`);

// Update configuration
const config = vscode.workspace.getConfiguration('superdesign');
await config.update('llmProvider', providerType, vscode.ConfigurationTarget.Global);

// Get the new provider
return await this.getProvider(providerType);
}

Expand All @@ -114,32 +106,37 @@ export class LLMProviderFactory {
type: LLMProviderType.CLAUDE_CODE,
name: 'Claude Code Binary',
description: 'Uses local claude-code binary for enhanced code execution capabilities'
}
},
{
type: LLMProviderType.CLAUDE_CODE_SDK,
name: 'Claude Code SDK',
description: 'Uses @anthropic-ai/claude-code SDK query() for in-process streaming without spawning a binary'
},
];
}

async validateProvider(providerType: LLMProviderType): Promise<{ isValid: boolean; error?: string }> {
try {
const provider = await this.createProvider(providerType);
const isValid = await provider.waitForInitialization();

if (!isValid) {
return { isValid: false, error: `Failed to initialize ${provider.getProviderName()}` };
}

// Additional validation based on provider type
if (!provider.hasValidConfiguration()) {
let errorMessage = '';

switch (providerType) {
case LLMProviderType.CLAUDE_API:
errorMessage = 'API key is required for Claude API provider';
break;
case LLMProviderType.CLAUDE_CODE:
errorMessage = 'Claude Code binary is not available. Please install claude-code CLI tool.';
break;
case LLMProviderType.CLAUDE_CODE_SDK:
errorMessage = 'Claude Code SDK is not available.';
break;
}

return { isValid: false, error: errorMessage };
}

Expand All @@ -150,7 +147,6 @@ export class LLMProviderFactory {
}
}

// Method to get provider status for UI display
async getProviderStatus(): Promise<{
current: LLMProviderType;
providers: Array<{
Expand All @@ -162,7 +158,7 @@ export class LLMProviderFactory {
}> {
const currentType = this.getConfiguredProviderType();
const availableProviders = this.getAvailableProviders();

const providerStatuses = await Promise.all(
availableProviders.map(async (provider) => {
const validation = await this.validateProvider(provider.type);
Expand All @@ -175,15 +171,11 @@ export class LLMProviderFactory {
})
);

return {
current: currentType,
providers: providerStatuses
};
return { current: currentType, providers: providerStatuses };
}

// Clean up all providers
dispose(): void {
this.providers.clear();
this.currentProvider = null;
}
}
}
4 changes: 2 additions & 2 deletions src/services/claudeCodeService.ts
Original file line number Diff line number Diff line change
Expand Up @@ -79,7 +79,7 @@ export class ClaudeCodeService {
}

// Method to get current provider information
async getProviderInfo(): Promise<{ name: string; type: 'api' | 'binary' }> {
async getProviderInfo(): Promise<{ name: string; type: 'api' | 'binary' | 'sdk' }> {
try {
const provider = await this.getCurrentProvider();
return {
Expand All @@ -88,7 +88,7 @@ export class ClaudeCodeService {
};
} catch (error) {
Logger.error(`Failed to get provider info: ${error}`);
return { name: 'Unknown', type: 'api' };
return { name: 'Unknown', type: 'api' as const };
}
}

Expand Down
Loading
Loading