Register configured local PromptKit backend
This commit is contained in:
@@ -171,6 +171,7 @@ func buildProductionLLMClient(ctx context.Context, cfg config.Config, profileID
|
|||||||
client, err := llm.NewPromptKitClient(llm.PromptKitClientConfig{
|
client, err := llm.NewPromptKitClient(llm.PromptKitClientConfig{
|
||||||
ProfileDir: cfg.PromptKit.ProfileDir,
|
ProfileDir: cfg.PromptKit.ProfileDir,
|
||||||
ProfileFile: cfg.PromptKit.ProfileFile,
|
ProfileFile: cfg.PromptKit.ProfileFile,
|
||||||
|
LocalBackend: mapPromptKitLocalBackend(cfg.PromptKit.LocalBackend),
|
||||||
Assets: assets,
|
Assets: assets,
|
||||||
Recorder: recorder,
|
Recorder: recorder,
|
||||||
ReasoningEffort: overrides.ReasoningEffort,
|
ReasoningEffort: overrides.ReasoningEffort,
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io/fs"
|
"io/fs"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"reflect"
|
"reflect"
|
||||||
@@ -13,7 +15,9 @@ import (
|
|||||||
"sort"
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
|
"testing/fstest"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"gitea.maximumdirect.net/eric/notarius/internal/core/artifacts"
|
"gitea.maximumdirect.net/eric/notarius/internal/core/artifacts"
|
||||||
@@ -34,6 +38,7 @@ import (
|
|||||||
itemeventnormalize "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/normalize/itemevents"
|
itemeventnormalize "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/normalize/itemevents"
|
||||||
spellnormalize "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/normalize/spells"
|
spellnormalize "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/normalize/spells"
|
||||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/generic/normalize/noop"
|
"gitea.maximumdirect.net/eric/notarius/internal/modules/generic/normalize/noop"
|
||||||
|
"gitea.maximumdirect.net/eric/promptkit"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestProductionCatalogCoversMaintainedConfigurations(t *testing.T) {
|
func TestProductionCatalogCoversMaintainedConfigurations(t *testing.T) {
|
||||||
@@ -406,6 +411,89 @@ func TestProductionLLMClientFactoriesBuildOfflineRuntime(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestProductionLLMClientFactoryUsesConfiguredLocalBackend(t *testing.T) {
|
||||||
|
var providerCalls atomic.Int32
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
providerCalls.Add(1)
|
||||||
|
if r.URL.Path != "/v1/chat/completions" {
|
||||||
|
t.Errorf("provider path = %q, want /v1/chat/completions", r.URL.Path)
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_, _ = w.Write([]byte(`{
|
||||||
|
"choices": [{"message": {"role": "assistant", "content": "{\"ok\":true}"}}],
|
||||||
|
"usage": {"prompt_tokens": 3, "completion_tokens": 4, "total_tokens": 7}
|
||||||
|
}`))
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
profilePath := filepath.Join(t.TempDir(), "profiles.yml")
|
||||||
|
if err := os.WriteFile(profilePath, []byte(`id: local-profile
|
||||||
|
backend: local
|
||||||
|
model: local-model
|
||||||
|
`), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
assets := llm.NewAssetRegistry()
|
||||||
|
if err := assets.RegisterPromptFS(fstest.MapFS{
|
||||||
|
"production.local.yaml": {Data: []byte(`id: production.local
|
||||||
|
version: "v1"
|
||||||
|
inputs:
|
||||||
|
- name: transcript
|
||||||
|
required: true
|
||||||
|
messages:
|
||||||
|
- role: user
|
||||||
|
content: '{{ input "transcript" }}'
|
||||||
|
output:
|
||||||
|
format: json
|
||||||
|
validation_mode: json
|
||||||
|
`)},
|
||||||
|
}, "."); err != nil {
|
||||||
|
t.Fatalf("register prompt assets: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := config.Default()
|
||||||
|
cfg.PromptKit.ProfileFile = profilePath
|
||||||
|
cfg.PromptKit.LocalBackend = &config.PromptKitLocalBackendConfig{
|
||||||
|
Endpoint: server.URL + "/v1",
|
||||||
|
ConcurrencyLimit: 2,
|
||||||
|
}
|
||||||
|
client, manifests, err := productionLLMClientFactoryWithAssets(assets)(
|
||||||
|
context.Background(),
|
||||||
|
cfg,
|
||||||
|
"local-profile",
|
||||||
|
LLMRuntimeOverrides{},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("build production LLM runtime: %v", err)
|
||||||
|
}
|
||||||
|
if len(manifests) != 0 {
|
||||||
|
t.Fatalf("eager profile manifests = %#v, want none", manifests)
|
||||||
|
}
|
||||||
|
|
||||||
|
var out map[string]any
|
||||||
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
||||||
|
PromptID: "production.local",
|
||||||
|
ProfileID: "local-profile",
|
||||||
|
Inputs: contracts.LLMInputSet{
|
||||||
|
"transcript": contracts.NewLLMInputMaterial("transcript", "text/plain", []byte("local request"), "", ""),
|
||||||
|
},
|
||||||
|
}, &out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CompleteStructured() error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
if providerCalls.Load() != 1 {
|
||||||
|
t.Fatalf("provider calls = %d, want 1", providerCalls.Load())
|
||||||
|
}
|
||||||
|
provider, ok := client.(contracts.LLMProfileManifestProvider)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("production client %T does not provide profile manifests", client)
|
||||||
|
}
|
||||||
|
recorded := provider.LLMProfileManifests()
|
||||||
|
if len(recorded) != 1 || recorded[0].BackendID != promptkit.BackendLocal {
|
||||||
|
t.Fatalf("production profile manifests = %#v, want local backend", recorded)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestProductionLLMClientFactoriesRejectInvalidConstruction(t *testing.T) {
|
func TestProductionLLMClientFactoriesRejectInvalidConstruction(t *testing.T) {
|
||||||
t.Run("canceled context", func(t *testing.T) {
|
t.Run("canceled context", func(t *testing.T) {
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"testing/fstest"
|
"testing/fstest"
|
||||||
|
|
||||||
"gitea.maximumdirect.net/eric/notarius/internal/core/config"
|
"gitea.maximumdirect.net/eric/notarius/internal/core/config"
|
||||||
|
"gitea.maximumdirect.net/eric/notarius/internal/framework/llm"
|
||||||
"gitea.maximumdirect.net/eric/promptkit"
|
"gitea.maximumdirect.net/eric/promptkit"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -61,8 +62,21 @@ func newProfileValidationEngine(cfg config.Config) (*promptkit.Engine, error) {
|
|||||||
if cfg.PromptKit.ProfileFile != "" {
|
if cfg.PromptKit.ProfileFile != "" {
|
||||||
opts = append(opts, promptkit.WithProfileFile(cfg.PromptKit.ProfileFile))
|
opts = append(opts, promptkit.WithProfileFile(cfg.PromptKit.ProfileFile))
|
||||||
}
|
}
|
||||||
|
if localBackend := mapPromptKitLocalBackend(cfg.PromptKit.LocalBackend); localBackend != nil {
|
||||||
|
opts = append(opts, llm.PromptKitLocalBackendOption(*localBackend))
|
||||||
|
}
|
||||||
return promptkit.NewEngine(promptkit.Config{
|
return promptkit.NewEngine(promptkit.Config{
|
||||||
PromptDir: "unused",
|
PromptDir: "unused",
|
||||||
ProfileDir: cfg.PromptKit.ProfileDir,
|
ProfileDir: cfg.PromptKit.ProfileDir,
|
||||||
}, opts...)
|
}, opts...)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func mapPromptKitLocalBackend(cfg *config.PromptKitLocalBackendConfig) *llm.PromptKitLocalBackendConfig {
|
||||||
|
if cfg == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return &llm.PromptKitLocalBackendConfig{
|
||||||
|
Endpoint: cfg.Endpoint,
|
||||||
|
ConcurrencyLimit: cfg.ConcurrencyLimit,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
54
internal/cli/promptkit_profiles_test.go
Normal file
54
internal/cli/promptkit_profiles_test.go
Normal file
@@ -0,0 +1,54 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"gitea.maximumdirect.net/eric/notarius/internal/core/config"
|
||||||
|
"gitea.maximumdirect.net/eric/promptkit"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestExplicitPromptKitProfileValidationUsesConfiguredLocalBackendWithoutGeneration(t *testing.T) {
|
||||||
|
var providerCalls atomic.Int32
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
|
||||||
|
providerCalls.Add(1)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
profilePath := filepath.Join(t.TempDir(), "profiles.yml")
|
||||||
|
if err := os.WriteFile(profilePath, []byte(`id: local-profile
|
||||||
|
backend: local
|
||||||
|
model: local-model
|
||||||
|
`), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
cfg := config.Default()
|
||||||
|
cfg.PromptKit.ProfileFile = profilePath
|
||||||
|
cfg.PromptKit.LocalBackend = &config.PromptKitLocalBackendConfig{
|
||||||
|
Endpoint: server.URL + "/v1",
|
||||||
|
ConcurrencyLimit: 2,
|
||||||
|
}
|
||||||
|
if err := validateExplicitPromptKitProfiles(context.Background(), cfg, []string{"local-profile"}); err != nil {
|
||||||
|
t.Fatalf("validateExplicitPromptKitProfiles() error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
if providerCalls.Load() != 0 {
|
||||||
|
t.Fatalf("provider calls during configured profile validation = %d, want 0", providerCalls.Load())
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg.PromptKit.LocalBackend = nil
|
||||||
|
err := validateExplicitPromptKitProfiles(context.Background(), cfg, []string{"local-profile"})
|
||||||
|
if err == nil ||
|
||||||
|
!strings.Contains(err.Error(), `validate PromptKit profile "local-profile"`) ||
|
||||||
|
!strings.Contains(err.Error(), promptkit.BackendLocal) {
|
||||||
|
t.Fatalf("validation without registration error = %v, want profile and local backend context", err)
|
||||||
|
}
|
||||||
|
if providerCalls.Load() != 0 {
|
||||||
|
t.Fatalf("provider calls after missing-registration validation = %d, want 0", providerCalls.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -19,9 +19,15 @@ import (
|
|||||||
|
|
||||||
const promptKitProviderName = "promptkit"
|
const promptKitProviderName = "promptkit"
|
||||||
|
|
||||||
|
type PromptKitLocalBackendConfig struct {
|
||||||
|
Endpoint string
|
||||||
|
ConcurrencyLimit int
|
||||||
|
}
|
||||||
|
|
||||||
type PromptKitClientConfig struct {
|
type PromptKitClientConfig struct {
|
||||||
ProfileDir string
|
ProfileDir string
|
||||||
ProfileFile string
|
ProfileFile string
|
||||||
|
LocalBackend *PromptKitLocalBackendConfig
|
||||||
Assets *AssetRegistry
|
Assets *AssetRegistry
|
||||||
Timeout time.Duration
|
Timeout time.Duration
|
||||||
HTTPClient *http.Client
|
HTTPClient *http.Client
|
||||||
@@ -46,6 +52,10 @@ type LLMProfileRecorder struct {
|
|||||||
var _ contracts.StructuredLLMClient = (*PromptKitClient)(nil)
|
var _ contracts.StructuredLLMClient = (*PromptKitClient)(nil)
|
||||||
var _ contracts.LLMProfileManifestProvider = (*PromptKitClient)(nil)
|
var _ contracts.LLMProfileManifestProvider = (*PromptKitClient)(nil)
|
||||||
|
|
||||||
|
func PromptKitLocalBackendOption(cfg PromptKitLocalBackendConfig) promptkit.Option {
|
||||||
|
return promptkit.WithBackend(promptkit.LocalBackend(cfg.Endpoint, cfg.ConcurrencyLimit))
|
||||||
|
}
|
||||||
|
|
||||||
func NewPromptKitClient(cfg PromptKitClientConfig) (*PromptKitClient, error) {
|
func NewPromptKitClient(cfg PromptKitClientConfig) (*PromptKitClient, error) {
|
||||||
if cfg.Assets == nil {
|
if cfg.Assets == nil {
|
||||||
return nil, fmt.Errorf("PromptKit client assets must not be nil")
|
return nil, fmt.Errorf("PromptKit client assets must not be nil")
|
||||||
@@ -60,6 +70,11 @@ func NewPromptKitClient(cfg PromptKitClientConfig) (*PromptKitClient, error) {
|
|||||||
if profileFile := strings.TrimSpace(cfg.ProfileFile); profileFile != "" {
|
if profileFile := strings.TrimSpace(cfg.ProfileFile); profileFile != "" {
|
||||||
options = append(options, promptkit.WithProfileFile(profileFile))
|
options = append(options, promptkit.WithProfileFile(profileFile))
|
||||||
}
|
}
|
||||||
|
if cfg.LocalBackend != nil {
|
||||||
|
localBackend := *cfg.LocalBackend
|
||||||
|
localBackend.Endpoint = strings.TrimSpace(localBackend.Endpoint)
|
||||||
|
options = append(options, PromptKitLocalBackendOption(localBackend))
|
||||||
|
}
|
||||||
options = append(options, cfg.EngineOptions...)
|
options = append(options, cfg.EngineOptions...)
|
||||||
|
|
||||||
engine, err := promptkit.NewEngine(promptkit.Config{
|
engine, err := promptkit.NewEngine(promptkit.Config{
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"reflect"
|
"reflect"
|
||||||
@@ -334,6 +335,88 @@ func TestPromptKitClientUsesPromptDefaultProfileWhenRequestProfileEmpty(t *testi
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestPromptKitClientUsesConfiguredLocalBackend(t *testing.T) {
|
||||||
|
var providerCalls atomic.Int32
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
providerCalls.Add(1)
|
||||||
|
if r.URL.Path != "/v1/chat/completions" {
|
||||||
|
t.Errorf("provider path = %q, want /v1/chat/completions", r.URL.Path)
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_, _ = w.Write([]byte(`{
|
||||||
|
"choices": [{"message": {"role": "assistant", "content": "{\"ok\":true}"}}],
|
||||||
|
"usage": {"prompt_tokens": 3, "completion_tokens": 4, "total_tokens": 7}
|
||||||
|
}`))
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
profilePath := filepath.Join(t.TempDir(), "profiles.yml")
|
||||||
|
if err := os.WriteFile(profilePath, []byte(`id: local-profile
|
||||||
|
backend: local
|
||||||
|
model: local-model
|
||||||
|
`), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
localBackend := &PromptKitLocalBackendConfig{
|
||||||
|
Endpoint: " " + server.URL + "/v1 ",
|
||||||
|
ConcurrencyLimit: 2,
|
||||||
|
}
|
||||||
|
client, err := NewPromptKitClient(PromptKitClientConfig{
|
||||||
|
Assets: newTestPromptKitAssets(t),
|
||||||
|
ProfileFile: profilePath,
|
||||||
|
LocalBackend: localBackend,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewPromptKitClient() error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
localBackend.Endpoint = "http://127.0.0.1:1/v1"
|
||||||
|
|
||||||
|
request := contracts.StructuredCompletionRequest{
|
||||||
|
PromptID: "adapter.direct-session",
|
||||||
|
ProfileID: "local-profile",
|
||||||
|
Inputs: contracts.LLMInputSet{
|
||||||
|
"transcript": contracts.NewLLMInputMaterial("transcript", "application/json", []byte(`{"source":true}`), "", ""),
|
||||||
|
},
|
||||||
|
Vars: map[string]any{"custom": "value"},
|
||||||
|
}
|
||||||
|
var out map[string]any
|
||||||
|
response, err := client.CompleteStructured(context.Background(), request, &out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CompleteStructured() error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
if providerCalls.Load() != 1 {
|
||||||
|
t.Fatalf("provider calls = %d, want 1", providerCalls.Load())
|
||||||
|
}
|
||||||
|
if response.ProfileID != "local-profile" || response.Model != "local-model" {
|
||||||
|
t.Fatalf("response metadata = %#v", response)
|
||||||
|
}
|
||||||
|
if response.Debug == nil || response.Debug.Prompt == nil ||
|
||||||
|
response.Debug.Prompt.SelectedBackendID != promptkit.BackendLocal {
|
||||||
|
t.Fatalf("response debug prompt = %#v, want local backend", response.Debug)
|
||||||
|
}
|
||||||
|
manifests := client.LLMProfileManifests()
|
||||||
|
if len(manifests) != 1 || manifests[0].BackendID != promptkit.BackendLocal {
|
||||||
|
t.Fatalf("profile manifests = %#v, want local backend", manifests)
|
||||||
|
}
|
||||||
|
|
||||||
|
clientWithoutRegistration, err := NewPromptKitClient(PromptKitClientConfig{
|
||||||
|
Assets: newTestPromptKitAssets(t),
|
||||||
|
ProfileFile: profilePath,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewPromptKitClient() without local registration error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
_, err = clientWithoutRegistration.CompleteStructured(context.Background(), request, &out)
|
||||||
|
if err == nil ||
|
||||||
|
!strings.Contains(err.Error(), "prepare PromptKit prompt") ||
|
||||||
|
!strings.Contains(err.Error(), promptkit.BackendLocal) {
|
||||||
|
t.Fatalf("CompleteStructured() without registration error = %v, want preparation failure with local backend context", err)
|
||||||
|
}
|
||||||
|
if providerCalls.Load() != 1 {
|
||||||
|
t.Fatalf("provider calls after missing-registration failure = %d, want 1", providerCalls.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestLLMProfileRecorderDistinguishesEffectiveTargets(t *testing.T) {
|
func TestLLMProfileRecorderDistinguishesEffectiveTargets(t *testing.T) {
|
||||||
recorder := NewLLMProfileRecorder()
|
recorder := NewLLMProfileRecorder()
|
||||||
for _, profile := range []artifacts.LLMProfileManifest{
|
for _, profile := range []artifacts.LLMProfileManifest{
|
||||||
|
|||||||
Reference in New Issue
Block a user