mirror of
https://github.com/ollama/ollama.git
synced 2026-09-11 05:33:32 -04:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6c4f69b1eb | ||
|
|
589354c87b |
No files matched your search
@@ -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 {
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
Reference in new issue
Block a user