diff --git a/src/providers/claudeCodeSdkProvider.ts b/src/providers/claudeCodeSdkProvider.ts new file mode 100644 index 00000000..9b1d8dd0 --- /dev/null +++ b/src/providers/claudeCodeSdkProvider.ts @@ -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 { + 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 { + 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 { + const config = vscode.workspace.getConfiguration('superdesign'); + const modelId = config.get('claudeCodeModelId'); + if (modelId) { + this.modelId = modelId; + } + } + + async query( + prompt: string, + options?: Partial, + abortController?: AbortController, + onMessage?: LLMStreamCallback + ): Promise { + 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 { + 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 { + 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'; + } +} diff --git a/src/providers/llmProvider.ts b/src/providers/llmProvider.ts index a96716ff..8834fb56 100644 --- a/src/providers/llmProvider.ts +++ b/src/providers/llmProvider.ts @@ -51,7 +51,7 @@ export abstract class LLMProvider { abstract refreshConfiguration(): Promise; abstract isAuthError(errorMessage: string): boolean; abstract getProviderName(): string; - abstract getProviderType(): 'api' | 'binary'; + abstract getProviderType(): 'api' | 'binary' | 'sdk'; protected async ensureInitialized(): Promise { if (this.initializationPromise) { @@ -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 { @@ -79,4 +80,4 @@ export interface LLMProviderConfig { claudeCodePath?: string; modelId?: string; thinkingBudgetTokens?: number; -} \ No newline at end of file +} diff --git a/src/providers/llmProviderFactory.ts b/src/providers/llmProviderFactory.ts index 89a0f44f..9c75c68a 100644 --- a/src/providers/llmProviderFactory.ts +++ b/src/providers/llmProviderFactory.ts @@ -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 { @@ -19,12 +20,10 @@ export class LLMProviderFactory { } async getProvider(providerType?: LLMProviderType): Promise { - // 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()) { @@ -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 { 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}`); } @@ -64,9 +60,10 @@ export class LLMProviderFactory { private getConfiguredProviderType(): LLMProviderType { const config = vscode.workspace.getConfiguration('superdesign'); const providerType = config.get('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': @@ -83,7 +80,6 @@ export class LLMProviderFactory { if (!this.currentProvider) { return false; } - try { return await this.currentProvider.refreshConfiguration(); } catch (error) { @@ -94,12 +90,8 @@ export class LLMProviderFactory { async switchProvider(providerType: LLMProviderType): Promise { 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); } @@ -114,7 +106,12 @@ 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' + }, ]; } @@ -122,15 +119,13 @@ export class LLMProviderFactory { 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'; @@ -138,8 +133,10 @@ export class LLMProviderFactory { 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 }; } @@ -150,7 +147,6 @@ export class LLMProviderFactory { } } - // Method to get provider status for UI display async getProviderStatus(): Promise<{ current: LLMProviderType; providers: Array<{ @@ -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); @@ -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; } -} \ No newline at end of file +} diff --git a/src/services/claudeCodeService.ts b/src/services/claudeCodeService.ts index 9fb98bbd..1c163e83 100644 --- a/src/services/claudeCodeService.ts +++ b/src/services/claudeCodeService.ts @@ -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 { @@ -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 }; } } diff --git a/src/test/llm-service.test.ts b/src/test/llm-service.test.ts index 8fa5e5e4..e3bca698 100644 --- a/src/test/llm-service.test.ts +++ b/src/test/llm-service.test.ts @@ -1,8 +1,63 @@ -// SPDX-License-Identifier: AGPL-3.0 -import { test } from 'node:test'; -import assert from 'node:assert'; +import { describe, it } from 'node:test'; +import assert from 'node:assert/strict'; -test('llm-service: placeholder passes', () => { - // TODO(Phase 1): test provider selection, API key resolution, streaming - assert.ok(true); +// STORY-001: CC SDK integration +// RED tests — ClaudeCodeSdkProvider does not exist yet; these will fail until GREEN. +// The test imports will throw at module load time, which counts as test failure. + +describe('ClaudeCodeSdkProvider', () => { + it('exports ClaudeCodeSdkProvider class', async () => { + const mod = await import('../providers/claudeCodeSdkProvider.js'); + assert.ok(typeof mod.ClaudeCodeSdkProvider === 'function', 'ClaudeCodeSdkProvider must be a class'); + }); + + it('getProviderType returns sdk', async () => { + const { ClaudeCodeSdkProvider } = await import('../providers/claudeCodeSdkProvider.js'); + const fakeChannel = { appendLine: () => {}, show: () => {}, hide: () => {}, dispose: () => {}, name: 'test' } as any; + const provider = new ClaudeCodeSdkProvider(fakeChannel); + assert.equal(provider.getProviderType(), 'sdk'); + }); + + it('getProviderName returns Claude Code SDK', async () => { + const { ClaudeCodeSdkProvider } = await import('../providers/claudeCodeSdkProvider.js'); + const fakeChannel = { appendLine: () => {}, show: () => {}, hide: () => {}, dispose: () => {}, name: 'test' } as any; + const provider = new ClaudeCodeSdkProvider(fakeChannel); + assert.equal(provider.getProviderName(), 'Claude Code SDK'); + }); + + it('hasValidConfiguration returns true (no API key required)', async () => { + const { ClaudeCodeSdkProvider } = await import('../providers/claudeCodeSdkProvider.js'); + const fakeChannel = { appendLine: () => {}, show: () => {}, hide: () => {}, dispose: () => {}, name: 'test' } as any; + const provider = new ClaudeCodeSdkProvider(fakeChannel); + assert.equal(provider.hasValidConfiguration(), true); + }); + + it('isAuthError detects auth-related messages', async () => { + const { ClaudeCodeSdkProvider } = await import('../providers/claudeCodeSdkProvider.js'); + const fakeChannel = { appendLine: () => {}, show: () => {}, hide: () => {}, dispose: () => {}, name: 'test' } as any; + const provider = new ClaudeCodeSdkProvider(fakeChannel); + assert.equal(provider.isAuthError('authentication failed'), true); + assert.equal(provider.isAuthError('random network error'), false); + }); +}); + +describe('LLMProviderType', () => { + it('includes CLAUDE_CODE_SDK variant', async () => { + const { LLMProviderType } = await import('../providers/llmProvider.js'); + assert.ok('CLAUDE_CODE_SDK' in LLMProviderType, 'LLMProviderType must have CLAUDE_CODE_SDK'); + assert.equal(LLMProviderType.CLAUDE_CODE_SDK, 'claude-code-sdk'); + }); +}); + +describe('LLMProviderFactory', () => { + it('createProvider handles claude-code-sdk type', async () => { + const { LLMProviderFactory } = await import('../providers/llmProviderFactory.js'); + const { LLMProviderType } = await import('../providers/llmProvider.js'); + const fakeChannel = { appendLine: () => {}, show: () => {}, hide: () => {}, dispose: () => {}, name: 'test' } as any; + const factory = LLMProviderFactory.getInstance(fakeChannel); + // Should not throw — provider instance is created + const provider = await (factory as any).createProvider(LLMProviderType.CLAUDE_CODE_SDK); + assert.ok(provider !== null); + assert.equal(provider.getProviderType(), 'sdk'); + }); });