mirror of
https://github.com/Syngnat/GoNavi.git
synced 2026-08-19 22:15:12 +08:00
- 新增 SessionChatProvider 接口,补齐非流式对话的会话态复用能力 - 为 Cursor Agent 和 CodeBuddy CLI 同步实现流式与非流式会话续接及状态持久化 - CustomProvider 补充会话态透传,统一 custom provider 的会话复用行为 - Service 新增 AIChatSendInSession,聊天主链路非流式回退改走带 session 的发送接口 - 保留原 AIChatSend 无状态语义,避免标题生成和记忆压缩污染主会话上下文 - 补充前后端定向测试,覆盖会话恢复、续接发送和前端回退分流
96 lines
2.7 KiB
Go
96 lines
2.7 KiB
Go
package provider
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"fmt"
|
||
"strings"
|
||
|
||
"GoNavi-Wails/internal/ai"
|
||
)
|
||
|
||
// CustomProvider 自定义 Provider,根据 apiFormat 选择底层协议
|
||
// 支持 openai / anthropic / gemini / cursor-agent 等 API 格式
|
||
type CustomProvider struct {
|
||
inner Provider
|
||
name string
|
||
}
|
||
|
||
// NewCustomProvider 创建自定义 Provider 实例
|
||
func NewCustomProvider(config ai.ProviderConfig) (Provider, error) {
|
||
// 根据 apiFormat 决定使用哪个底层协议,默认 openai
|
||
apiFormat := strings.ToLower(strings.TrimSpace(config.APIFormat))
|
||
if apiFormat == "" {
|
||
apiFormat = "openai"
|
||
}
|
||
if strings.TrimSpace(config.BaseURL) == "" && apiFormat != "claude-cli" && apiFormat != "codebuddy-cli" {
|
||
return nil, fmt.Errorf("自定义 Provider 必须指定 Base URL")
|
||
}
|
||
|
||
var innerProvider Provider
|
||
var err error
|
||
switch apiFormat {
|
||
case "anthropic":
|
||
innerProvider, err = NewAnthropicProvider(config)
|
||
case "gemini":
|
||
innerProvider, err = NewGeminiProvider(config)
|
||
case "cursor-agent":
|
||
innerProvider, err = NewCursorAgentProvider(config)
|
||
case "claude-cli":
|
||
innerProvider, err = NewClaudeCLIProvider(config)
|
||
case "codebuddy-cli":
|
||
innerProvider, err = NewCodeBuddyCLIProvider(config)
|
||
default: // "openai" 及其他
|
||
innerProvider, err = NewOpenAIProvider(config)
|
||
}
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
name := strings.TrimSpace(config.Name)
|
||
if name == "" {
|
||
name = "Custom"
|
||
}
|
||
|
||
return &CustomProvider{
|
||
inner: innerProvider,
|
||
name: name,
|
||
}, nil
|
||
}
|
||
|
||
func (p *CustomProvider) Name() string {
|
||
return p.name
|
||
}
|
||
|
||
func (p *CustomProvider) Validate() error {
|
||
if strings.TrimSpace(p.inner.(interface{ Name() string }).Name()) == "" {
|
||
// 对自定义 Provider,API Key 可选(部分本地服务不需要)
|
||
}
|
||
return nil
|
||
}
|
||
|
||
func (p *CustomProvider) Chat(ctx context.Context, req ai.ChatRequest) (*ai.ChatResponse, error) {
|
||
return p.inner.Chat(ctx, req)
|
||
}
|
||
|
||
func (p *CustomProvider) ChatStream(ctx context.Context, req ai.ChatRequest, callback func(ai.StreamChunk)) error {
|
||
return p.inner.ChatStream(ctx, req, callback)
|
||
}
|
||
|
||
func (p *CustomProvider) ChatWithState(ctx context.Context, state json.RawMessage, req ai.ChatRequest) (*ai.ChatResponse, json.RawMessage, error) {
|
||
sessionProvider, ok := p.inner.(SessionChatProvider)
|
||
if !ok {
|
||
resp, err := p.inner.Chat(ctx, req)
|
||
return resp, nil, err
|
||
}
|
||
return sessionProvider.ChatWithState(ctx, state, req)
|
||
}
|
||
|
||
func (p *CustomProvider) ChatStreamWithState(ctx context.Context, state json.RawMessage, req ai.ChatRequest, callback func(ai.StreamChunk)) (json.RawMessage, error) {
|
||
sessionProvider, ok := p.inner.(SessionStreamProvider)
|
||
if !ok {
|
||
return nil, p.inner.ChatStream(ctx, req, callback)
|
||
}
|
||
return sessionProvider.ChatStreamWithState(ctx, state, req, callback)
|
||
}
|