providers: support GitHub Copilot stdio transport
This commit is contained in:
parent
dd54601f2d
commit
9d71743165
3 changed files with 161 additions and 42 deletions
|
|
@ -318,13 +318,13 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
|
||||||
|
|
||||||
case "github-copilot", "copilot":
|
case "github-copilot", "copilot":
|
||||||
apiBase := cfg.APIBase
|
apiBase := cfg.APIBase
|
||||||
if apiBase == "" {
|
|
||||||
apiBase = "localhost:4321"
|
|
||||||
}
|
|
||||||
connectMode := cfg.ConnectMode
|
connectMode := cfg.ConnectMode
|
||||||
if connectMode == "" {
|
if connectMode == "" {
|
||||||
connectMode = "grpc"
|
connectMode = "grpc"
|
||||||
}
|
}
|
||||||
|
if connectMode == "grpc" && apiBase == "" {
|
||||||
|
apiBase = "localhost:4321"
|
||||||
|
}
|
||||||
provider, err := NewGitHubCopilotProvider(apiBase, connectMode, modelID)
|
provider, err := NewGitHubCopilotProvider(apiBase, connectMode, modelID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, "", err
|
return nil, "", err
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,8 @@ import (
|
||||||
type GitHubCopilotProvider struct {
|
type GitHubCopilotProvider struct {
|
||||||
uri string
|
uri string
|
||||||
connectMode string // "stdio" or "grpc"
|
connectMode string // "stdio" or "grpc"
|
||||||
|
model string
|
||||||
|
closed bool
|
||||||
|
|
||||||
client *copilot.Client
|
client *copilot.Client
|
||||||
session *copilot.Session
|
session *copilot.Session
|
||||||
|
|
@ -24,46 +26,94 @@ func NewGitHubCopilotProvider(uri string, connectMode string, model string) (*Gi
|
||||||
connectMode = "grpc"
|
connectMode = "grpc"
|
||||||
}
|
}
|
||||||
|
|
||||||
switch connectMode {
|
if connectMode != "stdio" && connectMode != "grpc" {
|
||||||
|
return nil, fmt.Errorf("unknown connect mode: %s", connectMode)
|
||||||
|
}
|
||||||
|
|
||||||
|
provider := &GitHubCopilotProvider{
|
||||||
|
uri: uri,
|
||||||
|
connectMode: connectMode,
|
||||||
|
model: model,
|
||||||
|
}
|
||||||
|
|
||||||
|
if connectMode == "stdio" {
|
||||||
|
return provider, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := provider.ensureSession(context.Background(), model); err != nil {
|
||||||
|
provider.Close()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return provider, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *GitHubCopilotProvider) newClient() (*copilot.Client, error) {
|
||||||
|
switch p.connectMode {
|
||||||
case "stdio":
|
case "stdio":
|
||||||
// TODO: Implement stdio mode for GitHub Copilot provider
|
return copilot.NewClient(nil), nil
|
||||||
// See https://github.com/github/copilot-sdk/blob/main/docs/getting-started.md for details
|
|
||||||
return nil, fmt.Errorf("stdio mode not implemented for GitHub Copilot provider; please use 'grpc' mode instead")
|
|
||||||
case "grpc":
|
case "grpc":
|
||||||
client := copilot.NewClient(&copilot.ClientOptions{
|
return copilot.NewClient(&copilot.ClientOptions{
|
||||||
CLIUrl: uri,
|
CLIUrl: p.uri,
|
||||||
})
|
}), nil
|
||||||
if err := client.Start(context.Background()); err != nil {
|
default:
|
||||||
|
return nil, fmt.Errorf("unknown connect mode: %s", p.connectMode)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *GitHubCopilotProvider) ensureSession(ctx context.Context, model string) (*copilot.Session, error) {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
|
||||||
|
if p.closed {
|
||||||
|
return nil, fmt.Errorf("provider closed")
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.session != nil {
|
||||||
|
return p.session, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.client == nil {
|
||||||
|
client, err := p.newClient()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
p.client = client
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := p.client.Start(ctx); err != nil {
|
||||||
|
p.client.Stop()
|
||||||
|
p.client = nil
|
||||||
return nil, fmt.Errorf(
|
return nil, fmt.Errorf(
|
||||||
"can't connect to Github Copilot: %w; `https://github.com/github/copilot-sdk/blob/main/docs/getting-started.md#connecting-to-an-external-cli-server` for details",
|
"can't connect to Github Copilot: %w; `https://github.com/github/copilot-sdk/blob/main/docs/getting-started.md#connecting-to-an-external-cli-server` for details",
|
||||||
err,
|
err,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
session, err := client.CreateSession(context.Background(), &copilot.SessionConfig{
|
sessionModel := p.model
|
||||||
Model: model,
|
if model != "" {
|
||||||
|
sessionModel = model
|
||||||
|
}
|
||||||
|
|
||||||
|
session, err := p.client.CreateSession(ctx, &copilot.SessionConfig{
|
||||||
|
Model: sessionModel,
|
||||||
OnPermissionRequest: copilot.PermissionHandler.ApproveAll,
|
OnPermissionRequest: copilot.PermissionHandler.ApproveAll,
|
||||||
Hooks: &copilot.SessionHooks{},
|
Hooks: &copilot.SessionHooks{},
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
client.Stop()
|
p.client.Stop()
|
||||||
|
p.client = nil
|
||||||
return nil, fmt.Errorf("create session failed: %w", err)
|
return nil, fmt.Errorf("create session failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return &GitHubCopilotProvider{
|
p.session = session
|
||||||
uri: uri,
|
return session, nil
|
||||||
connectMode: connectMode,
|
|
||||||
client: client,
|
|
||||||
session: session,
|
|
||||||
}, nil
|
|
||||||
default:
|
|
||||||
return nil, fmt.Errorf("unknown connect mode: %s", connectMode)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *GitHubCopilotProvider) Close() {
|
func (p *GitHubCopilotProvider) Close() {
|
||||||
p.mu.Lock()
|
p.mu.Lock()
|
||||||
defer p.mu.Unlock()
|
defer p.mu.Unlock()
|
||||||
|
p.closed = true
|
||||||
if p.client != nil {
|
if p.client != nil {
|
||||||
p.client.Stop()
|
p.client.Stop()
|
||||||
p.client = nil
|
p.client = nil
|
||||||
|
|
@ -94,12 +144,9 @@ func (p *GitHubCopilotProvider) Chat(
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("marshal messages: %w", err)
|
return nil, fmt.Errorf("marshal messages: %w", err)
|
||||||
}
|
}
|
||||||
p.mu.Lock()
|
session, err := p.ensureSession(ctx, model)
|
||||||
session := p.session
|
if err != nil {
|
||||||
p.mu.Unlock()
|
return nil, err
|
||||||
|
|
||||||
if session == nil {
|
|
||||||
return nil, fmt.Errorf("provider closed")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
resp, err := session.SendAndWait(ctx, copilot.MessageOptions{
|
resp, err := session.SendAndWait(ctx, copilot.MessageOptions{
|
||||||
|
|
|
||||||
72
pkg/providers/github_copilot_provider_test.go
Normal file
72
pkg/providers/github_copilot_provider_test.go
Normal file
|
|
@ -0,0 +1,72 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewGitHubCopilotProviderStdioIsLazy(t *testing.T) {
|
||||||
|
provider, err := NewGitHubCopilotProvider("", "stdio", "gpt-4.1")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewGitHubCopilotProvider() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider == nil {
|
||||||
|
t.Fatal("NewGitHubCopilotProvider() returned nil")
|
||||||
|
}
|
||||||
|
if provider.connectMode != "stdio" {
|
||||||
|
t.Fatalf("connectMode = %q, want stdio", provider.connectMode)
|
||||||
|
}
|
||||||
|
if provider.client != nil {
|
||||||
|
t.Fatalf("client = %#v, want nil before first Chat", provider.client)
|
||||||
|
}
|
||||||
|
if provider.session != nil {
|
||||||
|
t.Fatalf("session = %#v, want nil before first Chat", provider.session)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateProviderUsesStdioWithoutGrpcDefaultApiBase(t *testing.T) {
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.ModelName = "test-copilot"
|
||||||
|
cfg.ModelList = []*config.ModelConfig{{
|
||||||
|
ModelName: "test-copilot",
|
||||||
|
Model: "github-copilot/gpt-4.1",
|
||||||
|
ConnectMode: "stdio",
|
||||||
|
}}
|
||||||
|
|
||||||
|
provider, _, err := CreateProvider(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CreateProvider() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
githubCopilotProvider, ok := provider.(*GitHubCopilotProvider)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("provider type = %T, want *GitHubCopilotProvider", provider)
|
||||||
|
}
|
||||||
|
if githubCopilotProvider.uri != "" {
|
||||||
|
t.Fatalf("uri = %q, want empty for stdio mode", githubCopilotProvider.uri)
|
||||||
|
}
|
||||||
|
if githubCopilotProvider.connectMode != "stdio" {
|
||||||
|
t.Fatalf("connectMode = %q, want stdio", githubCopilotProvider.connectMode)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGitHubCopilotProviderClosePreventsReopen(t *testing.T) {
|
||||||
|
provider, err := NewGitHubCopilotProvider("", "stdio", "gpt-4.1")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewGitHubCopilotProvider() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
provider.Close()
|
||||||
|
|
||||||
|
_, err = provider.Chat(context.Background(), nil, nil, "gpt-4.1", nil)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("Chat() error = nil, want provider closed")
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "provider closed") {
|
||||||
|
t.Fatalf("Chat() error = %v, want provider closed", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue