Compare commits

...
Author SHA1 Message Date
Willie Tran 6c4f69b1eb app: keep cloud settings copy unchanged 2026-08-26 15:46:40 -05:00
Willie Tran 589354c87b app: add desktop feature flag service 2026-08-26 15:43:01 -05:00
9 changed files with 661 additions and 47 deletions

No files matched your search

+4
View File
@@ -30,6 +30,7 @@ import (
"github.com/ollama/ollama/app/ui"
"github.com/ollama/ollama/app/updater"
"github.com/ollama/ollama/app/version"
ollamaAuth "github.com/ollama/ollama/auth"
)
var (
@@ -180,6 +181,9 @@ func main() {
// on macOS, offer the user to create a symlink
// from /usr/local/bin/ollama to the app bundle
installSymlink()
if err := ollamaAuth.EnsureKeypair(io.Discard); err != nil {
slog.Warn("failed to ensure signing identity", "error", err)
}
var ln net.Listener
if devMode {
+41
View File
@@ -6,10 +6,51 @@ vi.mock("./lib/ollama-client", () => ({
import {
fetchConnectUrl,
getFeatureFlag,
getClaudeDesktopAvailableModels,
getIntegrationStatuses,
} from "./api";
describe("getFeatureFlag", () => {
afterEach(() => {
vi.unstubAllGlobals();
});
it("returns boolean and string values from the local app service", async () => {
const fetch = vi
.fn()
.mockResolvedValueOnce(new Response(JSON.stringify({ value: true })))
.mockResolvedValueOnce(
new Response(JSON.stringify({ value: "compact" })),
);
vi.stubGlobal("fetch", fetch);
await expect(getFeatureFlag("new-chat", false)).resolves.toBe(true);
await expect(getFeatureFlag("chat-layout", "standard")).resolves.toBe(
"compact",
);
expect(fetch).toHaveBeenNthCalledWith(
1,
"http://127.0.0.1:3001/api/v1/feature-flags/new-chat?type=boolean&default=false",
);
expect(fetch).toHaveBeenNthCalledWith(
2,
"http://127.0.0.1:3001/api/v1/feature-flags/chat-layout?type=string&default=standard",
);
});
it("returns the compiled fallback for unavailable or invalid responses", async () => {
const fetch = vi
.fn()
.mockRejectedValueOnce(new Error("offline"))
.mockResolvedValueOnce(new Response(JSON.stringify({ value: "wrong" })));
vi.stubGlobal("fetch", fetch);
await expect(getFeatureFlag("new-chat", false)).resolves.toBe(false);
await expect(getFeatureFlag("new-chat", true)).resolves.toBe(true);
});
});
describe("fetchConnectUrl", () => {
afterEach(() => {
vi.unstubAllGlobals();
+32
View File
@@ -33,6 +33,38 @@ export interface CloudStatusResponse {
source: CloudStatusSource;
}
export async function getFeatureFlag(
key: string,
defaultValue: boolean,
): Promise<boolean>;
export async function getFeatureFlag(
key: string,
defaultValue: string,
): Promise<string>;
export async function getFeatureFlag(
key: string,
defaultValue: boolean | string,
): Promise<boolean | string> {
const type = typeof defaultValue === "boolean" ? "boolean" : "string";
const query = new URLSearchParams({
type,
default: String(defaultValue),
});
try {
const response = await fetch(
`${API_BASE}/api/v1/feature-flags/${encodeURIComponent(key)}?${query}`,
);
if (!response.ok) return defaultValue;
const data = await response.json();
return typeof data.value === typeof defaultValue
? data.value
: defaultValue;
} catch {
return defaultValue;
}
}
export interface IntegrationStatus {
id: string;
name: string;
+197
View File
@@ -0,0 +1,197 @@
//go:build windows || darwin
package ui
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"net/http"
"net/url"
"strconv"
"sync"
)
const maxFeatureFlagResponseSize = 4096
var errFeatureFlagUnavailable = errors.New("feature flag unavailable")
type featureFlagResult struct {
ready chan struct{}
value any
}
type featureFlagService struct {
mu sync.Mutex
results map[string]*featureFlagResult
fetch func(context.Context, string) (any, error)
cloudDisabled func() (bool, error)
}
func newFeatureFlagService(
fetch func(context.Context, string) (any, error),
cloudDisabled func() (bool, error),
) *featureFlagService {
return &featureFlagService{
results: make(map[string]*featureFlagResult),
fetch: fetch,
cloudDisabled: cloudDisabled,
}
}
func (s *featureFlagService) resolve(ctx context.Context, key string, defaultValue any) any {
if !validFeatureFlagKey(key) {
return defaultValue
}
s.mu.Lock()
result, ok := s.results[key]
if ok {
s.mu.Unlock()
select {
case <-result.ready:
return result.value
case <-ctx.Done():
return defaultValue
}
}
result = &featureFlagResult{ready: make(chan struct{})}
s.results[key] = result
s.mu.Unlock()
result.value = defaultValue
disabled, err := s.cloudDisabled()
if err == nil && !disabled {
value, err := s.fetch(ctx, key)
if err == nil {
switch defaultValue.(type) {
case bool:
if _, ok := value.(bool); ok {
result.value = value
}
case string:
if _, ok := value.(string); ok {
result.value = value
}
}
}
}
close(result.ready)
return result.value
}
func validFeatureFlagKey(key string) bool {
if key == "" || len(key) > 128 {
return false
}
for _, r := range key {
switch {
case r >= 'a' && r <= 'z':
case r >= 'A' && r <= 'Z':
case r >= '0' && r <= '9':
case r == '-', r == '_', r == '.', r == ':':
default:
return false
}
}
return true
}
func (s *Server) featureFlagResolver() *featureFlagService {
s.featureFlagsMu.Lock()
defer s.featureFlagsMu.Unlock()
if s.featureFlags == nil {
s.featureFlags = newFeatureFlagService(
s.fetchFeatureFlag,
func() (bool, error) {
if s.Store == nil {
return false, errFeatureFlagUnavailable
}
return s.Store.CloudDisabled()
},
)
}
return s.featureFlags
}
// FeatureFlagBool returns one session-stable boolean value or defaultValue.
func (s *Server) FeatureFlagBool(ctx context.Context, key string, defaultValue bool) bool {
valueBool, ok := s.featureFlagResolver().resolve(ctx, key, defaultValue).(bool)
if !ok {
return defaultValue
}
return valueBool
}
// FeatureFlagString returns one session-stable string value or defaultValue.
func (s *Server) FeatureFlagString(ctx context.Context, key, defaultValue string) string {
valueString, ok := s.featureFlagResolver().resolve(ctx, key, defaultValue).(string)
if !ok {
return defaultValue
}
return valueString
}
func (s *Server) fetchFeatureFlag(ctx context.Context, key string) (any, error) {
if !validFeatureFlagKey(key) {
return nil, errFeatureFlagUnavailable
}
resp, err := s.doSelfSigned(ctx, http.MethodGet, "/api/app/feature-flags/"+url.PathEscape(key))
if err != nil {
return nil, errFeatureFlagUnavailable
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, errFeatureFlagUnavailable
}
body, err := io.ReadAll(io.LimitReader(resp.Body, maxFeatureFlagResponseSize+1))
if err != nil || len(body) > maxFeatureFlagResponseSize {
return nil, errFeatureFlagUnavailable
}
var response struct {
Value any `json:"value"`
}
decoder := json.NewDecoder(bytes.NewReader(body))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&response); err != nil {
return nil, errFeatureFlagUnavailable
}
if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) {
return nil, errFeatureFlagUnavailable
}
switch value := response.Value.(type) {
case bool, string:
return value, nil
}
return nil, errFeatureFlagUnavailable
}
func (s *Server) getFeatureFlag(w http.ResponseWriter, r *http.Request) error {
key := r.PathValue("key")
defaultValue := r.URL.Query().Get("default")
var value any
switch r.URL.Query().Get("type") {
case "boolean":
if defaultValue != "true" && defaultValue != "false" {
w.WriteHeader(http.StatusBadRequest)
return nil
}
fallback, _ := strconv.ParseBool(defaultValue)
value = s.FeatureFlagBool(r.Context(), key, fallback)
case "string":
value = s.FeatureFlagString(r.Context(), key, defaultValue)
default:
w.WriteHeader(http.StatusBadRequest)
return nil
}
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Content-Type", "application/json")
return json.NewEncoder(w).Encode(struct {
Value any `json:"value"`
}{Value: value})
}
+298
View File
@@ -0,0 +1,298 @@
//go:build windows || darwin
package ui
import (
"context"
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"testing"
"golang.org/x/crypto/ssh"
)
func TestFeatureFlagsTypedValues(t *testing.T) {
server := &Server{
featureFlags: newFeatureFlagService(
func(_ context.Context, key string) (any, error) {
switch key {
case "enabled":
return true, nil
case "mode":
return "compact", nil
case "wrong-type":
return "yes", nil
default:
return nil, errors.New("unavailable")
}
},
func() (bool, error) { return false, nil },
),
}
if got := server.FeatureFlagBool(t.Context(), "enabled", false); !got {
t.Fatal("FeatureFlagBool() = false, want true")
}
if got := server.FeatureFlagString(t.Context(), "mode", "standard"); got != "compact" {
t.Fatalf("FeatureFlagString() = %q, want compact", got)
}
if got := server.FeatureFlagBool(t.Context(), "wrong-type", false); got {
t.Fatal("FeatureFlagBool() accepted a string value")
}
if got := server.FeatureFlagString(t.Context(), "missing", "standard"); got != "standard" {
t.Fatalf("FeatureFlagString() = %q, want fallback", got)
}
}
func TestFeatureFlagsResolveOncePerSession(t *testing.T) {
var calls atomic.Int32
started := make(chan struct{})
release := make(chan struct{})
server := &Server{
featureFlags: newFeatureFlagService(
func(context.Context, string) (any, error) {
if calls.Add(1) == 1 {
close(started)
}
<-release
return true, nil
},
func() (bool, error) { return false, nil },
),
}
const callers = 20
results := make(chan bool, callers)
var wg sync.WaitGroup
for range callers {
wg.Add(1)
go func() {
defer wg.Done()
results <- server.FeatureFlagBool(t.Context(), "shared", false)
}()
}
<-started
close(release)
wg.Wait()
close(results)
for result := range results {
if !result {
t.Fatal("FeatureFlagBool() = false, want true")
}
}
if got := calls.Load(); got != 1 {
t.Fatalf("remote calls = %d, want 1", got)
}
}
func TestFeatureFlagsCacheFailures(t *testing.T) {
for _, tt := range []struct {
name string
cloudDisabled func() (bool, error)
}{
{name: "cloud off", cloudDisabled: func() (bool, error) { return true, nil }},
{name: "cloud status unavailable", cloudDisabled: func() (bool, error) { return false, errors.New("unavailable") }},
{name: "remote unavailable", cloudDisabled: func() (bool, error) { return false, nil }},
} {
t.Run(tt.name, func(t *testing.T) {
var calls atomic.Int32
server := &Server{
featureFlags: newFeatureFlagService(
func(context.Context, string) (any, error) {
calls.Add(1)
return nil, errors.New("unavailable")
},
tt.cloudDisabled,
),
}
if got := server.FeatureFlagBool(t.Context(), "enabled", true); !got {
t.Fatal("first call did not return its fallback")
}
if got := server.FeatureFlagBool(t.Context(), "enabled", false); !got {
t.Fatal("second call did not return the session-stable fallback")
}
wantCalls := int32(1)
if tt.name != "remote unavailable" {
wantCalls = 0
}
if got := calls.Load(); got != wantCalls {
t.Fatalf("remote calls = %d, want %d", got, wantCalls)
}
})
}
}
func TestFeatureFlagsRejectInvalidKeysWithoutFetching(t *testing.T) {
var calls atomic.Int32
server := &Server{
featureFlags: newFeatureFlagService(
func(context.Context, string) (any, error) {
calls.Add(1)
return true, nil
},
func() (bool, error) { return false, nil },
),
}
for _, key := range []string{"", "flags/all", "has space", strings.Repeat("x", 129)} {
if got := server.FeatureFlagBool(t.Context(), key, false); got {
t.Fatalf("FeatureFlagBool(%q) = true, want fallback", key)
}
}
if got := calls.Load(); got != 0 {
t.Fatalf("remote calls = %d, want 0", got)
}
}
func TestFeatureFlagLocalAPI(t *testing.T) {
server := &Server{
Dev: true,
featureFlags: newFeatureFlagService(
func(_ context.Context, key string) (any, error) {
if key == "enabled" {
return true, nil
}
return "compact", nil
},
func() (bool, error) { return false, nil },
),
}
for _, tt := range []struct {
path string
want any
}{
{path: "/api/v1/feature-flags/enabled?type=boolean&default=false", want: true},
{path: "/api/v1/feature-flags/mode?type=string&default=standard", want: "compact"},
} {
rr := httptest.NewRecorder()
server.Handler().ServeHTTP(rr, httptest.NewRequest(http.MethodGet, tt.path, nil))
if rr.Code != http.StatusOK {
t.Fatalf("%s status = %d, want %d", tt.path, rr.Code, http.StatusOK)
}
var response struct {
Value any `json:"value"`
}
if err := json.NewDecoder(rr.Body).Decode(&response); err != nil {
t.Fatal(err)
}
if response.Value != tt.want {
t.Fatalf("%s value = %v, want %v", tt.path, response.Value, tt.want)
}
}
}
func TestFetchFeatureFlag(t *testing.T) {
home := t.TempDir()
t.Setenv("HOME", home)
writeFeatureFlagTestKey(t, home)
remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet || r.URL.Path != "/api/app/feature-flags/enabled" {
t.Fatalf("request = %s %s, want signed feature lookup", r.Method, r.URL.Path)
}
verifyFeatureFlagRequest(t, r)
w.Header().Set("Content-Type", "application/json")
fmt.Fprint(w, `{"value":true}`)
}))
defer remote.Close()
previous := OllamaDotCom
OllamaDotCom = remote.URL
t.Cleanup(func() { OllamaDotCom = previous })
value, err := (&Server{}).fetchFeatureFlag(t.Context(), "enabled")
if err != nil {
t.Fatal(err)
}
if value != true {
t.Fatalf("value = %v, want true", value)
}
}
func TestFetchFeatureFlagRejectsUnexpectedResponses(t *testing.T) {
home := t.TempDir()
t.Setenv("HOME", home)
writeFeatureFlagTestKey(t, home)
for _, body := range []string{
`{"value":null}`,
`{"value":1}`,
`{"value":{"enabled":true}}`,
`{"value":true,"reason":"forced"}`,
`{"value":true}{"value":false}`,
} {
t.Run(body, func(t *testing.T) {
remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
fmt.Fprint(w, body)
}))
defer remote.Close()
previous := OllamaDotCom
OllamaDotCom = remote.URL
defer func() { OllamaDotCom = previous }()
if _, err := (&Server{}).fetchFeatureFlag(t.Context(), "enabled"); err == nil {
t.Fatal("fetchFeatureFlag() accepted an unexpected response")
}
})
}
}
func writeFeatureFlagTestKey(t *testing.T, home string) {
t.Helper()
_, privateKey, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
t.Fatal(err)
}
block, err := ssh.MarshalPrivateKey(privateKey, "")
if err != nil {
t.Fatal(err)
}
keyPath := filepath.Join(home, ".ollama", "id_ed25519")
if err := os.MkdirAll(filepath.Dir(keyPath), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(keyPath, pem.EncodeToMemory(block), 0o600); err != nil {
t.Fatal(err)
}
}
func verifyFeatureFlagRequest(t *testing.T, req *http.Request) {
t.Helper()
keyData, signatureData, ok := strings.Cut(req.Header.Get("Authorization"), ":")
if !ok {
t.Fatal("request is missing its public-key signature")
}
keyData = strings.TrimPrefix(keyData, "Bearer ")
publicKeyData, err := base64.StdEncoding.DecodeString(keyData)
if err != nil {
t.Fatal(err)
}
publicKey, err := ssh.ParsePublicKey(publicKeyData)
if err != nil {
t.Fatal(err)
}
signature, err := base64.StdEncoding.DecodeString(signatureData)
if err != nil {
t.Fatal(err)
}
challenge := []byte(req.Method + "," + req.URL.RequestURI())
if err := publicKey.Verify(challenge, &ssh.Signature{Format: publicKey.Type(), Blob: signature}); err != nil {
t.Fatal(err)
}
}
+3
View File
@@ -115,6 +115,8 @@ type Server struct {
UpdateAvailableFunc func()
IntegrationInstalled func(string) bool
ListCloudModels func(context.Context) (*api.ListResponse, error)
featureFlagsMu sync.Mutex
featureFlags *featureFlagService
}
func (s *Server) log() *slog.Logger {
@@ -295,6 +297,7 @@ func (s *Server) Handler() http.Handler {
mux.Handle("POST /api/v1/settings", handle(s.settings))
mux.Handle("GET /api/v1/cloud", handle(s.getCloudSetting))
mux.Handle("POST /api/v1/cloud", handle(s.cloudSetting))
mux.Handle("GET /api/v1/feature-flags/{key}", handle(s.getFeatureFlag))
mux.Handle("GET /api/v1/models/cloud", handle(s.getCloudModels))
mux.Handle("GET /api/v1/integrations", handle(s.getIntegrationStatuses))
+44
View File
@@ -3,8 +3,10 @@ package auth
import (
"bytes"
"context"
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/pem"
"errors"
"fmt"
"io"
@@ -18,6 +20,48 @@ import (
const defaultPrivateKey = "id_ed25519"
// EnsureKeypair creates the default signing keypair when it does not exist.
func EnsureKeypair(out io.Writer) error {
home, err := os.UserHomeDir()
if err != nil {
return err
}
privKeyPath := filepath.Join(home, ".ollama", defaultPrivateKey)
pubKeyPath := privKeyPath + ".pub"
if _, err := os.Stat(privKeyPath); !os.IsNotExist(err) {
return nil
}
fmt.Fprintf(out, "Couldn't find '%s'. Generating new private key.\n", privKeyPath)
cryptoPublicKey, cryptoPrivateKey, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
return err
}
privateKeyBytes, err := ssh.MarshalPrivateKey(cryptoPrivateKey, "")
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(privKeyPath), 0o755); err != nil {
return fmt.Errorf("could not create directory %w", err)
}
if err := os.WriteFile(privKeyPath, pem.EncodeToMemory(privateKeyBytes), 0o600); err != nil {
return err
}
sshPublicKey, err := ssh.NewPublicKey(cryptoPublicKey)
if err != nil {
return err
}
publicKeyBytes := ssh.MarshalAuthorizedKey(sshPublicKey)
if err := os.WriteFile(pubKeyPath, publicKeyBytes, 0o644); err != nil {
return err
}
fmt.Fprintf(out, "Your new public key is: \n\n%s\n", publicKeyBytes)
return nil
}
func GetPublicKey() (string, error) {
home, err := os.UserHomeDir()
if err != nil {
+40
View File
@@ -0,0 +1,40 @@
package auth
import (
"bytes"
"context"
"os"
"path/filepath"
"testing"
)
func TestEnsureKeypair(t *testing.T) {
home := t.TempDir()
t.Setenv("HOME", home)
t.Setenv("USERPROFILE", home)
var output bytes.Buffer
if err := EnsureKeypair(&output); err != nil {
t.Fatal(err)
}
if output.Len() == 0 {
t.Fatal("EnsureKeypair() did not report the generated public key")
}
for _, name := range []string{"id_ed25519", "id_ed25519.pub"} {
if _, err := os.Stat(filepath.Join(home, ".ollama", name)); err != nil {
t.Fatalf("generated key %s: %v", name, err)
}
}
if _, err := Sign(context.Background(), []byte("request")); err != nil {
t.Fatalf("Sign() after EnsureKeypair(): %v", err)
}
output.Reset()
if err := EnsureKeypair(&output); err != nil {
t.Fatal(err)
}
if output.Len() != 0 {
t.Fatal("EnsureKeypair() replaced an existing key")
}
}
+2 -47
View File
@@ -3,10 +3,7 @@ package cmd
import (
"bufio"
"context"
"crypto/ed25519"
"crypto/rand"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"io"
@@ -33,11 +30,11 @@ import (
"github.com/olekukonko/tablewriter"
"github.com/pkg/browser"
"github.com/spf13/cobra"
"golang.org/x/crypto/ssh"
"golang.org/x/sync/errgroup"
"golang.org/x/term"
"github.com/ollama/ollama/api"
"github.com/ollama/ollama/auth"
"github.com/ollama/ollama/cmd/config"
"github.com/ollama/ollama/cmd/launch"
"github.com/ollama/ollama/cmd/tui"
@@ -2031,49 +2028,7 @@ func RunServer(_ *cobra.Command, _ []string) error {
}
func initializeKeypair() error {
home, err := os.UserHomeDir()
if err != nil {
return err
}
privKeyPath := filepath.Join(home, ".ollama", "id_ed25519")
pubKeyPath := filepath.Join(home, ".ollama", "id_ed25519.pub")
_, err = os.Stat(privKeyPath)
if os.IsNotExist(err) {
fmt.Printf("Couldn't find '%s'. Generating new private key.\n", privKeyPath)
cryptoPublicKey, cryptoPrivateKey, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
return err
}
privateKeyBytes, err := ssh.MarshalPrivateKey(cryptoPrivateKey, "")
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(privKeyPath), 0o755); err != nil {
return fmt.Errorf("could not create directory %w", err)
}
if err := os.WriteFile(privKeyPath, pem.EncodeToMemory(privateKeyBytes), 0o600); err != nil {
return err
}
sshPublicKey, err := ssh.NewPublicKey(cryptoPublicKey)
if err != nil {
return err
}
publicKeyBytes := ssh.MarshalAuthorizedKey(sshPublicKey)
if err := os.WriteFile(pubKeyPath, publicKeyBytes, 0o644); err != nil {
return err
}
fmt.Printf("Your new public key is: \n\n%s\n", publicKeyBytes)
}
return nil
return auth.EnsureKeypair(os.Stdout)
}
func checkServerHeartbeat(cmd *cobra.Command, _ []string) error {