Skip to content
Open
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
59 changes: 59 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
name: CI

on:
pull_request:
push:
branches:
- master

permissions:
contents: read

jobs:
backend:
name: Go tests and vet
runs-on: ubuntu-22.04
steps:
- name: Checkout
uses: actions/checkout@v4

- name: Set up Go
uses: actions/setup-go@v5
with:
go-version-file: go.mod
cache: true

- name: Install Wails build dependencies
run: sudo apt-get update && sudo apt-get install -y libgtk-3-dev libwebkit2gtk-4.1-dev

- name: Test
run: go test -tags webkit2_41 ./...

- name: Vet
run: go vet -tags webkit2_41 ./...

frontend:
name: Frontend tests and build
runs-on: ubuntu-latest
defaults:
run:
working-directory: frontend
steps:
- name: Checkout
uses: actions/checkout@v4

- name: Set up Node
uses: actions/setup-node@v4
with:
node-version: 20
cache: npm
cache-dependency-path: frontend/package-lock.json

- name: Install dependencies
run: npm ci

- name: Test
run: npm test

- name: Build
run: npm run build
66 changes: 53 additions & 13 deletions app.go
Original file line number Diff line number Diff line change
Expand Up @@ -231,9 +231,14 @@ func (a *App) GetAppState() (config.AppState, error) {
defaultProviderID := firstNonEmpty(piDefaults.DefaultProvider, cfg.Settings.LastDefaultProviderID, selectedProvider)
defaultModelID := firstNonEmpty(piDefaults.DefaultModel, cfg.Settings.LastDefaultModelID)

providerTransports, err := provider.ConfigTransports(cfg.Providers)
if err != nil {
return config.AppState{}, err
}

return config.AppState{
Version: appVersion,
Providers: cfg.Providers,
Providers: providerTransports,
SelectedProviderID: selectedProvider,
DefaultProviderID: defaultProviderID,
DefaultModelID: defaultModelID,
Expand All @@ -242,20 +247,39 @@ func (a *App) GetAppState() (config.AppState, error) {
}, nil
}

func (a *App) ListProviders() ([]provider.Config, error) {
func (a *App) ListProviders() ([]provider.ConfigTransport, error) {
cfg, err := a.coordinator.Load()
if err != nil {
return nil, err
}
return cfg.Providers, nil
return provider.ConfigTransports(cfg.Providers)
}

func (a *App) CreateProvider(input provider.Config) error {
return a.coordinator.UpsertProvider("", input)
func (a *App) CreateProvider(input provider.ConfigTransport) (provider.ConfigTransport, error) {
converted, err := input.Config()
if err != nil {
return provider.ConfigTransport{}, err
}
if err := a.coordinator.UpsertProvider("", converted); err != nil {
return provider.ConfigTransport{}, err
}
cfg, err := a.coordinator.Load()
if err != nil {
return provider.ConfigTransport{}, err
}
persisted, err := cfg.ProviderByID(converted.ID)
if err != nil {
return provider.ConfigTransport{}, err
}
return provider.NewConfigTransport(persisted)
}

func (a *App) UpdateProvider(id string, input provider.Config) error {
return a.coordinator.UpsertProvider(id, input)
func (a *App) UpdateProvider(id string, input provider.ConfigTransport) error {
converted, err := input.Config()
if err != nil {
return err
}
return a.coordinator.UpsertProvider(id, converted)
}

func (a *App) DeleteProvider(id string) error {
Expand Down Expand Up @@ -318,7 +342,7 @@ func (a *App) TestConnection(id string) (provider.ConnectionTestResult, error) {
}, nil
}

func (a *App) FetchModels(id string) ([]provider.ModelInfo, error) {
func (a *App) FetchModels(id string) ([]provider.ModelTransport, error) {
cfg, err := a.coordinator.Load()
if err != nil {
return nil, err
Expand All @@ -331,16 +355,32 @@ func (a *App) FetchModels(id string) ([]provider.ModelInfo, error) {
if current.APIKeyEnv != "" && !envResult.Found {
return nil, errors.New("环境变量 " + current.APIKeyEnv + " 不存在")
}
return provider.FetchModelsByAPI(current, key)
models, err := provider.FetchModelsByAPI(current, key)
if err != nil {
return nil, err
}
return provider.ModelTransports(models)
}

func (a *App) ImportModels(providerID string, models []provider.ModelInfo) error {
return a.coordinator.MergeModels(providerID, models)
func (a *App) ImportModels(providerID string, models []provider.ModelTransport) error {
converted, err := provider.ModelsFromTransport(models)
if err != nil {
return err
}
return a.coordinator.MergeModels(providerID, converted)
}

// ReplaceModels 用给定列表整体替换该 provider 的模型集合(替换语义,未传入的将被删除)。
func (a *App) ReplaceModels(providerID string, models []provider.ModelInfo) error {
return a.coordinator.ReplaceModels(providerID, models)
func (a *App) ReplaceModels(providerID string, models []provider.ModelTransport, expectedRevision string) (provider.ModelListTransport, error) {
converted, err := provider.ModelsFromTransport(models)
if err != nil {
return provider.ModelListTransport{}, err
}
replaced, err := a.coordinator.ReplaceModels(providerID, converted, expectedRevision)
if err != nil {
return provider.ModelListTransport{}, err
}
return provider.NewModelListTransport(replaced)
}

func (a *App) SetDefaultModel(providerID string, modelID string) error {
Expand Down
1 change: 1 addition & 0 deletions frontend/package.json
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
"scripts": {
"dev": "vite",
"build": "vite build",
"test": "node --test src/config/model-editor.test.js src/actions/provider-actions.test.js",
"preview": "vite preview"
},
"devDependencies": {
Expand Down
Loading