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":
|
||||
apiBase := cfg.APIBase
|
||||
if apiBase == "" {
|
||||
apiBase = "localhost:4321"
|
||||
}
|
||||
connectMode := cfg.ConnectMode
|
||||
if connectMode == "" {
|
||||
connectMode = "grpc"
|
||||
}
|
||||
if connectMode == "grpc" && apiBase == "" {
|
||||
apiBase = "localhost:4321"
|
||||
}
|
||||
provider, err := NewGitHubCopilotProvider(apiBase, connectMode, modelID)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
|
|
|
|||
|
|
@ -12,6 +12,8 @@ import (
|
|||
type GitHubCopilotProvider struct {
|
||||
uri string
|
||||
connectMode string // "stdio" or "grpc"
|
||||
model string
|
||||
closed bool
|
||||
|
||||
client *copilot.Client
|
||||
session *copilot.Session
|
||||
|
|
@ -24,46 +26,94 @@ func NewGitHubCopilotProvider(uri string, connectMode string, model string) (*Gi
|
|||
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":
|
||||
// TODO: Implement stdio mode for GitHub Copilot provider
|
||||
// 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")
|
||||
return copilot.NewClient(nil), nil
|
||||
case "grpc":
|
||||
client := copilot.NewClient(&copilot.ClientOptions{
|
||||
CLIUrl: uri,
|
||||
})
|
||||
if err := client.Start(context.Background()); err != nil {
|
||||
return copilot.NewClient(&copilot.ClientOptions{
|
||||
CLIUrl: p.uri,
|
||||
}), 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(
|
||||
"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,
|
||||
)
|
||||
}
|
||||
|
||||
session, err := client.CreateSession(context.Background(), &copilot.SessionConfig{
|
||||
Model: model,
|
||||
sessionModel := p.model
|
||||
if model != "" {
|
||||
sessionModel = model
|
||||
}
|
||||
|
||||
session, err := p.client.CreateSession(ctx, &copilot.SessionConfig{
|
||||
Model: sessionModel,
|
||||
OnPermissionRequest: copilot.PermissionHandler.ApproveAll,
|
||||
Hooks: &copilot.SessionHooks{},
|
||||
})
|
||||
if err != nil {
|
||||
client.Stop()
|
||||
p.client.Stop()
|
||||
p.client = nil
|
||||
return nil, fmt.Errorf("create session failed: %w", err)
|
||||
}
|
||||
|
||||
return &GitHubCopilotProvider{
|
||||
uri: uri,
|
||||
connectMode: connectMode,
|
||||
client: client,
|
||||
session: session,
|
||||
}, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unknown connect mode: %s", connectMode)
|
||||
}
|
||||
p.session = session
|
||||
return session, nil
|
||||
}
|
||||
|
||||
func (p *GitHubCopilotProvider) Close() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.closed = true
|
||||
if p.client != nil {
|
||||
p.client.Stop()
|
||||
p.client = nil
|
||||
|
|
@ -94,12 +144,9 @@ func (p *GitHubCopilotProvider) Chat(
|
|||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal messages: %w", err)
|
||||
}
|
||||
p.mu.Lock()
|
||||
session := p.session
|
||||
p.mu.Unlock()
|
||||
|
||||
if session == nil {
|
||||
return nil, fmt.Errorf("provider closed")
|
||||
session, err := p.ensureSession(ctx, model)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
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