mirror of
https://github.com/mudler/LocalAI.git
synced 2026-10-01 02:24:34 -04:00
feat(localai-proxy): add a backend that serves text APIs from a remote LocalAI
Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
1 parent
509cca35e0
commit
6d5c9600d7
14 files changed
+1621
-4
No files matched your search
@@ -5948,6 +5948,35 @@ include:
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
# localai-proxy
|
||||
- build-type: ''
|
||||
cuda-major-version: ""
|
||||
cuda-minor-version: ""
|
||||
platforms: 'linux/amd64'
|
||||
platform-tag: 'amd64'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-cpu-localai-proxy'
|
||||
runs-on: 'ubuntu-latest'
|
||||
base-image: "ubuntu:24.04"
|
||||
skip-drivers: 'false'
|
||||
backend: "localai-proxy"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: ''
|
||||
cuda-major-version: ""
|
||||
cuda-minor-version: ""
|
||||
platforms: 'linux/arm64'
|
||||
platform-tag: 'arm64'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-cpu-localai-proxy'
|
||||
runs-on: 'ubuntu-24.04-arm'
|
||||
base-image: "ubuntu:24.04"
|
||||
skip-drivers: 'false'
|
||||
backend: "localai-proxy"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
# valkey-store
|
||||
- build-type: ''
|
||||
cuda-major-version: ""
|
||||
@@ -6753,6 +6782,10 @@ includeDarwin:
|
||||
tag-suffix: "-metal-darwin-arm64-cloud-proxy"
|
||||
build-type: "metal"
|
||||
lang: "go"
|
||||
- backend: "localai-proxy"
|
||||
tag-suffix: "-metal-darwin-arm64-localai-proxy"
|
||||
build-type: "metal"
|
||||
lang: "go"
|
||||
- backend: "valkey-store"
|
||||
tag-suffix: "-metal-darwin-arm64-valkey-store"
|
||||
build-type: "metal"
|
||||
|
||||
@@ -29,6 +29,7 @@ LocalAI
|
||||
# Root-level build artifacts when running `go build ./...` against
|
||||
# Go backend packages whose main lives under backend/go/.
|
||||
/cloud-proxy
|
||||
/localai-proxy
|
||||
/local-store
|
||||
/valkey-store
|
||||
# prevent above rules from omitting the helm chart
|
||||
@@ -50,6 +51,8 @@ tests/e2e-aio/backends
|
||||
/tests/e2e/mock-backend/mock-backend
|
||||
# The cloud-proxy backend binary the e2e suite runs next to the mock backend.
|
||||
/tests/e2e/mock-backend/cloud-proxy
|
||||
# The localai-proxy backend binary, built next to it the same way.
|
||||
/tests/e2e/mock-backend/localai-proxy
|
||||
|
||||
release/
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
# Disable parallel execution for backend builds
|
||||
.NOTPARALLEL: backends/diffusers backends/llama-cpp backends/turboquant backends/bonsai backends/outetts backends/piper backends/stablediffusion-ggml backends/trellis2cpp backends/trellis2cpp-darwin backends/whisper backends/crispasr backends/parakeet-cpp backends/moss-transcribe-cpp backends/nemo-speech-cpp backends/faster-whisper backends/silero-vad backends/local-store backends/valkey-store backends/cloud-proxy backends/huggingface backends/rfdetr backends/rfdetr-cpp backends/insightface backends/speaker-recognition backends/kitten-tts backends/kokoro backends/chatterbox backends/llama-cpp-darwin backends/neutts build-darwin-python-backend build-darwin-go-backend backends/mlx backends/mlx-video backends/diffuser-darwin backends/mlx-vlm backends/mlx-audio backends/mlx-distributed backends/stablediffusion-ggml-darwin backends/vllm backends/vllm-omni backends/longcat-video backends/sglang backends/moonshine backends/pocket-tts backends/qwen-tts backends/faster-qwen3-tts backends/qwen-asr backends/nemo backends/voxcpm backends/whisperx backends/ace-step backends/acestep-cpp backends/fish-speech backends/voxtral backends/opus backends/trl backends/llama-cpp-quantization backends/kokoros backends/sam3-cpp backends/qwen3-tts-cpp backends/moss-tts-cpp backends/magpie-tts-cpp backends/vllm-cpp backends/omnivoice-cpp backends/vibevoice-cpp backends/localvqe backends/tinygrad backends/sherpa-onnx backends/ds4 backends/ds4-darwin backends/liquid-audio backends/supertonic backends/depth-anything-cpp backends/privacy-filter backends/privacy-filter-darwin backends/audio-cpp backends/audio-cpp-darwin
|
||||
.NOTPARALLEL: backends/diffusers backends/llama-cpp backends/turboquant backends/bonsai backends/outetts backends/piper backends/stablediffusion-ggml backends/trellis2cpp backends/trellis2cpp-darwin backends/whisper backends/crispasr backends/parakeet-cpp backends/moss-transcribe-cpp backends/nemo-speech-cpp backends/faster-whisper backends/silero-vad backends/local-store backends/valkey-store backends/cloud-proxy backends/localai-proxy backends/huggingface backends/rfdetr backends/rfdetr-cpp backends/insightface backends/speaker-recognition backends/kitten-tts backends/kokoro backends/chatterbox backends/llama-cpp-darwin backends/neutts build-darwin-python-backend build-darwin-go-backend backends/mlx backends/mlx-video backends/diffuser-darwin backends/mlx-vlm backends/mlx-audio backends/mlx-distributed backends/stablediffusion-ggml-darwin backends/vllm backends/vllm-omni backends/longcat-video backends/sglang backends/moonshine backends/pocket-tts backends/qwen-tts backends/faster-qwen3-tts backends/qwen-asr backends/nemo backends/voxcpm backends/whisperx backends/ace-step backends/acestep-cpp backends/fish-speech backends/voxtral backends/opus backends/trl backends/llama-cpp-quantization backends/kokoros backends/sam3-cpp backends/qwen3-tts-cpp backends/moss-tts-cpp backends/magpie-tts-cpp backends/vllm-cpp backends/omnivoice-cpp backends/vibevoice-cpp backends/localvqe backends/tinygrad backends/sherpa-onnx backends/ds4 backends/ds4-darwin backends/liquid-audio backends/supertonic backends/depth-anything-cpp backends/privacy-filter backends/privacy-filter-darwin backends/audio-cpp backends/audio-cpp-darwin
|
||||
.NOTPARALLEL: backends/whisper-medusa
|
||||
.NOTPARALLEL: backends/funasr
|
||||
|
||||
@@ -76,7 +76,7 @@ else
|
||||
GORELEASER=$(shell which goreleaser)
|
||||
endif
|
||||
|
||||
TEST_PATHS?=./api/... ./pkg/... ./core/... ./backend/go/cloud-proxy/... ./backend/go/local-store/... ./backend/go/valkey-store/...
|
||||
TEST_PATHS?=./api/... ./pkg/... ./core/... ./backend/go/cloud-proxy/... ./backend/go/localai-proxy/... ./backend/go/local-store/... ./backend/go/valkey-store/...
|
||||
|
||||
## Coverage output and the committed baseline that CI compares against.
|
||||
## The gate is strict: total coverage must never decrease (no tolerance).
|
||||
@@ -385,13 +385,14 @@ prepare-e2e:
|
||||
run-e2e-image:
|
||||
docker run -p 5390:8080 -e MODELS_PATH=/models -e THREADS=1 -e DEBUG=true -d --rm -v $(TEST_DIR):/models --name e2e-tests-$(RANDOM) localai-tests
|
||||
|
||||
test-e2e: build-mock-backend build-cloud-proxy-backend prepare-e2e run-e2e-image
|
||||
test-e2e: build-mock-backend build-cloud-proxy-backend build-localai-proxy-backend prepare-e2e run-e2e-image
|
||||
@echo 'Running e2e tests'
|
||||
BUILD_TYPE=$(BUILD_TYPE) \
|
||||
LOCALAI_API=http://$(E2E_BRIDGE_IP):5390 \
|
||||
$(GOCMD) run github.com/onsi/ginkgo/v2/ginkgo --flake-attempts $(TEST_FLAKES) -v -r ./tests/e2e
|
||||
$(MAKE) clean-mock-backend
|
||||
$(MAKE) clean-cloud-proxy-backend
|
||||
$(MAKE) clean-localai-proxy-backend
|
||||
$(MAKE) teardown-e2e
|
||||
docker rmi localai-tests
|
||||
|
||||
@@ -1330,6 +1331,7 @@ BACKEND_PIPER = piper|golang|.|false|true
|
||||
BACKEND_LOCAL_STORE = local-store|golang|.|false|true
|
||||
BACKEND_VALKEY_STORE = valkey-store|golang|.|false|true
|
||||
BACKEND_CLOUD_PROXY = cloud-proxy|golang|.|false|true
|
||||
BACKEND_LOCALAI_PROXY = localai-proxy|golang|.|false|true
|
||||
BACKEND_HUGGINGFACE = huggingface|golang|.|false|true
|
||||
BACKEND_SILERO_VAD = silero-vad|golang|.|false|true
|
||||
BACKEND_STABLEDIFFUSION_GGML = stablediffusion-ggml|golang|.|--progress=plain|true
|
||||
@@ -1436,6 +1438,7 @@ $(eval $(call generate-docker-build-target,$(BACKEND_PIPER)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_LOCAL_STORE)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_VALKEY_STORE)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_CLOUD_PROXY)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_LOCALAI_PROXY)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_HUGGINGFACE)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_SILERO_VAD)))
|
||||
$(eval $(call generate-docker-build-target,$(BACKEND_STABLEDIFFUSION_GGML)))
|
||||
@@ -1511,7 +1514,7 @@ $(eval $(call generate-docker-build-target,$(BACKEND_SUPERTONIC)))
|
||||
docker-save-%: backend-images
|
||||
docker save local-ai-backend:$* -o backend-images/$*.tar
|
||||
|
||||
docker-build-backends: docker-build-llama-cpp docker-build-ik-llama-cpp docker-build-turboquant docker-build-bonsai docker-build-ds4 docker-build-rerankers docker-build-vllm docker-build-vllm-omni docker-build-longcat-video docker-build-sglang docker-build-transformers docker-build-outetts docker-build-diffusers docker-build-kokoro docker-build-faster-whisper docker-build-crispasr docker-build-coqui docker-build-chatterbox docker-build-vibevoice docker-build-liquid-audio docker-build-moonshine docker-build-pocket-tts docker-build-qwen-tts docker-build-fish-speech docker-build-faster-qwen3-tts docker-build-qwen-asr docker-build-nemo docker-build-voxcpm docker-build-whisperx docker-build-ace-step docker-build-acestep-cpp docker-build-voxtral docker-build-mlx-distributed docker-build-trl docker-build-llama-cpp-quantization docker-build-tinygrad docker-build-kokoros docker-build-sam3-cpp docker-build-rfdetr-cpp docker-build-qwen3-tts-cpp docker-build-moss-tts-cpp docker-build-magpie-tts-cpp docker-build-vllm-cpp docker-build-omnivoice-cpp docker-build-vibevoice-cpp docker-build-localvqe docker-build-insightface docker-build-speaker-recognition docker-build-sherpa-onnx docker-build-cloud-proxy docker-build-supertonic docker-build-depth-anything-cpp docker-build-moss-transcribe-cpp docker-build-nemo-speech-cpp docker-build-privacy-filter docker-build-trellis2cpp docker-build-valkey-store docker-build-audio-cpp
|
||||
docker-build-backends: docker-build-llama-cpp docker-build-ik-llama-cpp docker-build-turboquant docker-build-bonsai docker-build-ds4 docker-build-rerankers docker-build-vllm docker-build-vllm-omni docker-build-longcat-video docker-build-sglang docker-build-transformers docker-build-outetts docker-build-diffusers docker-build-kokoro docker-build-faster-whisper docker-build-crispasr docker-build-coqui docker-build-chatterbox docker-build-vibevoice docker-build-liquid-audio docker-build-moonshine docker-build-pocket-tts docker-build-qwen-tts docker-build-fish-speech docker-build-faster-qwen3-tts docker-build-qwen-asr docker-build-nemo docker-build-voxcpm docker-build-whisperx docker-build-ace-step docker-build-acestep-cpp docker-build-voxtral docker-build-mlx-distributed docker-build-trl docker-build-llama-cpp-quantization docker-build-tinygrad docker-build-kokoros docker-build-sam3-cpp docker-build-rfdetr-cpp docker-build-qwen3-tts-cpp docker-build-moss-tts-cpp docker-build-magpie-tts-cpp docker-build-vllm-cpp docker-build-omnivoice-cpp docker-build-vibevoice-cpp docker-build-localvqe docker-build-insightface docker-build-speaker-recognition docker-build-sherpa-onnx docker-build-cloud-proxy docker-build-localai-proxy docker-build-supertonic docker-build-depth-anything-cpp docker-build-moss-transcribe-cpp docker-build-nemo-speech-cpp docker-build-privacy-filter docker-build-trellis2cpp docker-build-valkey-store docker-build-audio-cpp
|
||||
docker-build-backends: docker-build-whisper-medusa
|
||||
docker-build-backends: docker-build-funasr
|
||||
|
||||
@@ -1531,6 +1534,12 @@ build-cloud-proxy-backend: protogen-go
|
||||
clean-cloud-proxy-backend:
|
||||
rm -f tests/e2e/mock-backend/cloud-proxy
|
||||
|
||||
build-localai-proxy-backend: protogen-go
|
||||
$(GOCMD) build -o tests/e2e/mock-backend/localai-proxy ./backend/go/localai-proxy
|
||||
|
||||
clean-localai-proxy-backend:
|
||||
rm -f tests/e2e/mock-backend/localai-proxy
|
||||
|
||||
########################################################
|
||||
### UI E2E Test Server
|
||||
########################################################
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
GOCMD=go
|
||||
|
||||
# Packaged as a standalone gallery backend by backend/Dockerfile.golang.
|
||||
localai-proxy:
|
||||
CGO_ENABLED=0 $(GOCMD) build -ldflags "$(LD_FLAGS)" -tags "$(GO_TAGS)" -o localai-proxy ./
|
||||
|
||||
package:
|
||||
bash package.sh
|
||||
|
||||
build: localai-proxy package
|
||||
|
||||
clean:
|
||||
rm -f localai-proxy
|
||||
@@ -0,0 +1,220 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/xlog"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// maxErrorBody caps the upstream body quoted in an error. It keeps gRPC
|
||||
// status messages small while leaving room for LocalAI's JSON error text,
|
||||
// which failover scans for request errors such as context overflows.
|
||||
const maxErrorBody = 500
|
||||
|
||||
// postJSON sends body as JSON to path and decodes a 2xx JSON reply into out
|
||||
// (skipped when out is nil). The request_timeout_seconds limit applies.
|
||||
func (p *LocalAIProxy) postJSON(ctx context.Context, path string, body, out any) error {
|
||||
cfg, err := p.config()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
payload, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return status.Errorf(codes.InvalidArgument, "localai-proxy: encode %s request: %v", path, err)
|
||||
}
|
||||
ctx, cancel := withTimeout(ctx, cfg)
|
||||
defer cancel()
|
||||
|
||||
req, err := p.newRequest(ctx, cfg, path, bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
return p.do(req, path, out)
|
||||
}
|
||||
|
||||
// postMultipart sends fields and, when fileField is set, the local file at
|
||||
// filePath as a multipart form to path, decoding a 2xx JSON reply into out
|
||||
// (skipped when out is nil). Core hands audio and images to backends as local
|
||||
// paths, and LocalAI's upload endpoints take them as multipart files. The
|
||||
// request_timeout_seconds limit applies.
|
||||
func (p *LocalAIProxy) postMultipart(ctx context.Context, path string, fields map[string]string, fileField, filePath string, out any) error {
|
||||
cfg, err := p.config()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var file *os.File
|
||||
if fileField != "" {
|
||||
// Open before contacting the upstream so a bad path is reported as a
|
||||
// request error, not as a failure of the remote host.
|
||||
if file, err = os.Open(filePath); err != nil {
|
||||
return status.Errorf(codes.InvalidArgument, "localai-proxy: open %s: %v", filePath, err)
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
}
|
||||
ctx, cancel := withTimeout(ctx, cfg)
|
||||
defer cancel()
|
||||
|
||||
// Stream the form through a pipe so large audio files are not buffered
|
||||
// in memory. The transport closes pr when the request ends, which
|
||||
// unblocks the writer on every error path.
|
||||
pr, pw := io.Pipe()
|
||||
mw := multipart.NewWriter(pw)
|
||||
go func() {
|
||||
pw.CloseWithError(writeMultipart(mw, fields, fileField, file))
|
||||
}()
|
||||
|
||||
req, err := p.newRequest(ctx, cfg, path, pr)
|
||||
if err != nil {
|
||||
_ = pr.CloseWithError(err)
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", mw.FormDataContentType())
|
||||
return p.do(req, path, out)
|
||||
}
|
||||
|
||||
func writeMultipart(mw *multipart.Writer, fields map[string]string, fileField string, file *os.File) error {
|
||||
for k, v := range fields {
|
||||
if err := mw.WriteField(k, v); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if file != nil {
|
||||
part, err := mw.CreateFormFile(fileField, filepath.Base(file.Name()))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := io.Copy(part, file); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return mw.Close()
|
||||
}
|
||||
|
||||
// postStream sends body as JSON to path and returns the open response of a
|
||||
// 2xx reply; the caller must close its body. No request_timeout_seconds
|
||||
// limit applies: streams legitimately outlast it, and ctx bounds them.
|
||||
func (p *LocalAIProxy) postStream(ctx context.Context, path string, body any) (*http.Response, error) {
|
||||
cfg, err := p.config()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
payload, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "localai-proxy: encode %s request: %v", path, err)
|
||||
}
|
||||
req, err := p.newRequest(ctx, cfg, path, bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp, err := p.client.Do(req)
|
||||
if err != nil {
|
||||
return nil, transportError(path, err)
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode > 299 {
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
return nil, statusError(path, resp)
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) newRequest(ctx context.Context, cfg *proxyConfig, path string, body io.Reader) (*http.Request, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.base+path, body)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "localai-proxy: build %s request: %v", path, err)
|
||||
}
|
||||
if cfg.apiKey != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+cfg.apiKey)
|
||||
}
|
||||
return req, nil
|
||||
}
|
||||
|
||||
// do runs req and decodes a 2xx JSON reply into out.
|
||||
func (p *LocalAIProxy) do(req *http.Request, path string, out any) error {
|
||||
resp, err := p.client.Do(req)
|
||||
if err != nil {
|
||||
return transportError(path, err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode < 200 || resp.StatusCode > 299 {
|
||||
return statusError(path, resp)
|
||||
}
|
||||
if out == nil {
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
return nil
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(out); err != nil {
|
||||
if ctxErr := req.Context().Err(); ctxErr != nil {
|
||||
return transportError(path, ctxErr)
|
||||
}
|
||||
return status.Errorf(codes.Internal, "localai-proxy: decode %s response: %v", path, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func withTimeout(ctx context.Context, cfg *proxyConfig) (context.Context, context.CancelFunc) {
|
||||
if cfg.timeout > 0 {
|
||||
return context.WithTimeout(ctx, cfg.timeout)
|
||||
}
|
||||
return context.WithCancel(ctx)
|
||||
}
|
||||
|
||||
// transportError maps a failed round trip to a gRPC status. A dead or
|
||||
// unreachable upstream is Unavailable so failover moves to the next target.
|
||||
func transportError(path string, err error) error {
|
||||
code := codes.Unavailable
|
||||
switch {
|
||||
case errors.Is(err, context.DeadlineExceeded):
|
||||
code = codes.DeadlineExceeded
|
||||
case errors.Is(err, context.Canceled):
|
||||
code = codes.Canceled
|
||||
}
|
||||
xlog.Warn("localai-proxy: upstream request failed", "path", path, "error", err)
|
||||
return status.Errorf(code, "localai-proxy: upstream %s: %v", path, err)
|
||||
}
|
||||
|
||||
// statusError maps a non-2xx upstream reply to a gRPC status. 5xx means the
|
||||
// upstream is unhealthy (Unavailable, so failover retries elsewhere); 4xx
|
||||
// means the request itself is wrong (InvalidArgument, so failover does not
|
||||
// trip a healthy target over a client error). 501 is the upstream saying it
|
||||
// cannot serve this kind of request, which failover treats as a capability
|
||||
// gap, like our own Unimplemented methods.
|
||||
func statusError(path string, resp *http.Response) error {
|
||||
raw, _ := io.ReadAll(io.LimitReader(resp.Body, maxErrorBody+1))
|
||||
msg := strings.TrimSpace(string(raw))
|
||||
if len(msg) > maxErrorBody {
|
||||
msg = msg[:maxErrorBody] + "..."
|
||||
}
|
||||
// gRPC refuses to send a status message that is not valid UTF-8, and the
|
||||
// cut above may split a rune.
|
||||
msg = strings.ToValidUTF8(msg, "")
|
||||
|
||||
var code codes.Code
|
||||
switch {
|
||||
case resp.StatusCode == http.StatusNotImplemented:
|
||||
code = codes.Unimplemented
|
||||
case resp.StatusCode >= 500:
|
||||
code = codes.Unavailable
|
||||
case resp.StatusCode >= 400:
|
||||
code = codes.InvalidArgument
|
||||
default:
|
||||
// A 1xx/3xx here means a misbehaving upstream (redirects are refused
|
||||
// by the client), not a bad request.
|
||||
code = codes.Unavailable
|
||||
}
|
||||
xlog.Warn("localai-proxy: upstream error", "path", path, "status", resp.StatusCode)
|
||||
return status.Error(code, fmt.Sprintf("localai-proxy: upstream %s returned %d: %s", path, resp.StatusCode, msg))
|
||||
}
|
||||
@@ -0,0 +1,160 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"io"
|
||||
"mime"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
. "github.com/onsi/gomega"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
// recordedRequest is what the fake upstream saw for one call. JSON bodies land
|
||||
// in JSON; multipart bodies land in Fields and Files (field name to content).
|
||||
type recordedRequest struct {
|
||||
Method string
|
||||
Path string
|
||||
Auth string
|
||||
JSON map[string]any
|
||||
Fields map[string]string
|
||||
Files map[string]string
|
||||
}
|
||||
|
||||
// scriptedResponse is the reply for one path. SSE, when set, is written as
|
||||
// "data: <frame>" events and wins over Body.
|
||||
type scriptedResponse struct {
|
||||
Status int
|
||||
ContentType string
|
||||
Body string
|
||||
SSE []string
|
||||
}
|
||||
|
||||
// fakeUpstream stands in for a remote LocalAI: it records every request and
|
||||
// answers each path with the response scripted for it (404 otherwise).
|
||||
type fakeUpstream struct {
|
||||
*httptest.Server
|
||||
|
||||
mu sync.Mutex
|
||||
requests []recordedRequest
|
||||
responses map[string]scriptedResponse
|
||||
}
|
||||
|
||||
func newFakeUpstream() *fakeUpstream {
|
||||
f := &fakeUpstream{responses: map[string]scriptedResponse{}}
|
||||
f.Server = httptest.NewServer(http.HandlerFunc(f.serve))
|
||||
return f
|
||||
}
|
||||
|
||||
// newFakeUpstreamWithHandler serves every request with h instead of the
|
||||
// scripted responses, for tests that need to control timing.
|
||||
func newFakeUpstreamWithHandler(h http.HandlerFunc) *fakeUpstream {
|
||||
return &fakeUpstream{Server: httptest.NewServer(h), responses: map[string]scriptedResponse{}}
|
||||
}
|
||||
|
||||
func (f *fakeUpstream) script(path string, r scriptedResponse) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.responses[path] = r
|
||||
}
|
||||
|
||||
// replyJSON scripts a 200 JSON response for path.
|
||||
func (f *fakeUpstream) replyJSON(path string, body any) {
|
||||
raw, err := json.Marshal(body)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
f.script(path, scriptedResponse{Status: http.StatusOK, ContentType: "application/json", Body: string(raw)})
|
||||
}
|
||||
|
||||
func (f *fakeUpstream) recorded() []recordedRequest {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return append([]recordedRequest(nil), f.requests...)
|
||||
}
|
||||
|
||||
// last returns the single most recent request, failing when none arrived.
|
||||
func (f *fakeUpstream) last() recordedRequest {
|
||||
reqs := f.recorded()
|
||||
ExpectWithOffset(1, reqs).NotTo(BeEmpty(), "upstream received no request")
|
||||
return reqs[len(reqs)-1]
|
||||
}
|
||||
|
||||
func (f *fakeUpstream) serve(w http.ResponseWriter, r *http.Request) {
|
||||
rec := recordedRequest{Method: r.Method, Path: r.URL.Path, Auth: r.Header.Get("Authorization")}
|
||||
mediaType, params, _ := mime.ParseMediaType(r.Header.Get("Content-Type"))
|
||||
switch {
|
||||
case mediaType == "multipart/form-data":
|
||||
rec.Fields, rec.Files = map[string]string{}, map[string]string{}
|
||||
mr := multipart.NewReader(r.Body, params["boundary"])
|
||||
for {
|
||||
part, err := mr.NextPart()
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
data, _ := io.ReadAll(part)
|
||||
if part.FileName() != "" {
|
||||
rec.Files[part.FormName()] = string(data)
|
||||
} else {
|
||||
rec.Fields[part.FormName()] = string(data)
|
||||
}
|
||||
}
|
||||
default:
|
||||
raw, _ := io.ReadAll(r.Body)
|
||||
if len(raw) > 0 {
|
||||
_ = json.Unmarshal(raw, &rec.JSON)
|
||||
}
|
||||
}
|
||||
|
||||
f.mu.Lock()
|
||||
f.requests = append(f.requests, rec)
|
||||
resp, ok := f.responses[r.URL.Path]
|
||||
f.mu.Unlock()
|
||||
|
||||
if !ok {
|
||||
http.Error(w, "no scripted response for "+r.URL.Path, http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
if resp.SSE != nil {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
flusher, _ := w.(http.Flusher)
|
||||
for _, frame := range resp.SSE {
|
||||
_, _ = io.WriteString(w, "data: "+frame+"\n\n")
|
||||
if flusher != nil {
|
||||
flusher.Flush()
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
if resp.ContentType != "" {
|
||||
w.Header().Set("Content-Type", resp.ContentType)
|
||||
}
|
||||
w.WriteHeader(resp.Status)
|
||||
_, _ = io.WriteString(w, resp.Body)
|
||||
}
|
||||
|
||||
// loadProxy returns a proxy loaded against the fake upstream with the given
|
||||
// proxy options merged over sane defaults.
|
||||
func loadProxy(f *fakeUpstream, mutate func(*pb.ModelOptions)) *LocalAIProxy {
|
||||
opts := &pb.ModelOptions{
|
||||
Model: "local-name",
|
||||
Proxy: &pb.ProxyOptions{UpstreamUrl: f.URL + "/", UpstreamModel: "remote-model"},
|
||||
}
|
||||
if mutate != nil {
|
||||
mutate(opts)
|
||||
}
|
||||
p := NewLocalAIProxy()
|
||||
ExpectWithOffset(1, p.Load(opts)).To(Succeed())
|
||||
return p
|
||||
}
|
||||
|
||||
// sseJSON marshals v for use as one SSE frame.
|
||||
func sseJSON(v any) string {
|
||||
raw, err := json.Marshal(v)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
return strings.TrimSpace(string(raw))
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mudler/xlog"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
func TestLocalAIProxy(t *testing.T) {
|
||||
RegisterFailHandler(Fail)
|
||||
// The specs drive upstream failures on purpose; their warnings are
|
||||
// expected and would only bury real failures in the output.
|
||||
xlog.SetLogger(xlog.NewLogger(xlog.LogLevelError, xlog.TextFormat))
|
||||
RunSpecs(t, "localai-proxy specs")
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package main
|
||||
|
||||
// localai-proxy is a LocalAI backend that serves backend gRPC methods by
|
||||
// calling the REST API of another LocalAI instance. It lets a model config
|
||||
// (and so a failover chain target) live on a remote LocalAI while callers
|
||||
// keep using the local backend interface for every modality, not only chat.
|
||||
|
||||
import (
|
||||
"flag"
|
||||
"os"
|
||||
|
||||
grpc "github.com/mudler/LocalAI/pkg/grpc"
|
||||
"github.com/mudler/xlog"
|
||||
"golang.org/x/term"
|
||||
)
|
||||
|
||||
var addr = flag.String("addr", "localhost:50051", "the address to listen on")
|
||||
|
||||
func main() {
|
||||
// xlog's default handler emits ANSI color codes, which are unreadable once
|
||||
// LocalAI captures the backend's stdout into a log file. Force plain text
|
||||
// when LOCALAI_LOG_FORMAT is unset and stdout is not a terminal.
|
||||
format := os.Getenv("LOCALAI_LOG_FORMAT")
|
||||
if format == "" && !term.IsTerminal(int(os.Stdout.Fd())) {
|
||||
format = xlog.TextFormat
|
||||
}
|
||||
xlog.SetLogger(xlog.NewLogger(xlog.LogLevel(os.Getenv("LOCALAI_LOG_LEVEL")), format))
|
||||
flag.Parse()
|
||||
if err := grpc.StartServer(*addr, NewLocalAIProxy()); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
Executable
+13
@@ -0,0 +1,13 @@
|
||||
#!/bin/bash
|
||||
|
||||
# Script to copy the localai-proxy binary into the package dir for the
|
||||
# final Dockerfile stage. Mirrors backend/go/local-store/package.sh —
|
||||
# no extra runtime libs needed since the backend is pure Go.
|
||||
|
||||
set -e
|
||||
|
||||
CURDIR=$(dirname "$(realpath $0)")
|
||||
|
||||
mkdir -p $CURDIR/package
|
||||
cp -avf $CURDIR/localai-proxy $CURDIR/package/
|
||||
cp -rfv $CURDIR/run.sh $CURDIR/package/
|
||||
@@ -0,0 +1,226 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/xlog"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/mudler/LocalAI/pkg/grpc/base"
|
||||
"github.com/mudler/LocalAI/pkg/grpc/grpcerrors"
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/LocalAI/pkg/httpclient"
|
||||
)
|
||||
|
||||
const (
|
||||
backendName = "localai-proxy"
|
||||
|
||||
// realtimePipelineOption names the upstream realtime pipeline that serves
|
||||
// live transcription sessions (options: ["realtime_pipeline:<name>"]).
|
||||
realtimePipelineOption = "realtime_pipeline:"
|
||||
)
|
||||
|
||||
// LocalAIProxy serves backend methods by calling a remote LocalAI's REST API.
|
||||
// base.SingleThread is not embedded: every call is an independent HTTP
|
||||
// request, so serialising them would only add latency.
|
||||
type LocalAIProxy struct {
|
||||
base.Base
|
||||
|
||||
cfg atomic.Pointer[proxyConfig]
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
type proxyConfig struct {
|
||||
base string // upstream base URL without a trailing slash
|
||||
upstreamModel string // model name sent upstream
|
||||
apiKey string
|
||||
realtimePipeline string
|
||||
timeout time.Duration // per-request limit for non-streaming calls; 0 = none
|
||||
}
|
||||
|
||||
func NewLocalAIProxy() *LocalAIProxy {
|
||||
// httpclient.New refuses redirects: the upstream is one configured
|
||||
// LocalAI, so a 3xx means misconfiguration or a hijacked host, and
|
||||
// following it would replay the bearer key to an unvetted host. It also
|
||||
// sets no body deadline, so long SSE streams are not cut short.
|
||||
return &LocalAIProxy{client: httpclient.New()}
|
||||
}
|
||||
|
||||
// Load refuses a model without proxy options so greedy backend probing,
|
||||
// which tries every installed backend on a model file, never selects it.
|
||||
func (p *LocalAIProxy) Load(opts *pb.ModelOptions) error {
|
||||
po := opts.GetProxy()
|
||||
if po == nil {
|
||||
return errors.New("localai-proxy: Load requires proxy options (proxy.upstream_url)")
|
||||
}
|
||||
raw := po.GetUpstreamUrl()
|
||||
if raw == "" {
|
||||
return errors.New("localai-proxy: proxy.upstream_url is required")
|
||||
}
|
||||
u, err := url.ParseRequestURI(raw)
|
||||
if err != nil {
|
||||
return fmt.Errorf("localai-proxy: proxy.upstream_url %q invalid: %w", raw, err)
|
||||
}
|
||||
if (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" {
|
||||
return fmt.Errorf("localai-proxy: proxy.upstream_url %q must be an http(s) URL with a host", raw)
|
||||
}
|
||||
|
||||
// There is no translate mode: the upstream always speaks LocalAI's API.
|
||||
if po.GetMode() != "" || po.GetProvider() != "" {
|
||||
xlog.Warn("localai-proxy: proxy.mode and proxy.provider are ignored",
|
||||
"mode", po.GetMode(), "provider", po.GetProvider())
|
||||
}
|
||||
|
||||
key, err := resolveAPIKey(po.GetApiKeyEnv(), po.GetApiKeyFile())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
model := po.GetUpstreamModel()
|
||||
if model == "" {
|
||||
model = opts.GetModel()
|
||||
}
|
||||
if model == "" {
|
||||
xlog.Warn("localai-proxy: no upstream model name; set proxy.upstream_model")
|
||||
}
|
||||
|
||||
var pipeline string
|
||||
for _, o := range opts.GetOptions() {
|
||||
if v, ok := strings.CutPrefix(o, realtimePipelineOption); ok {
|
||||
pipeline = strings.TrimSpace(v)
|
||||
}
|
||||
}
|
||||
|
||||
var timeout time.Duration
|
||||
if s := po.GetRequestTimeoutSeconds(); s > 0 {
|
||||
timeout = time.Duration(s) * time.Second
|
||||
}
|
||||
|
||||
p.cfg.Store(&proxyConfig{
|
||||
base: strings.TrimRight(raw, "/"),
|
||||
upstreamModel: model,
|
||||
apiKey: key,
|
||||
realtimePipeline: pipeline,
|
||||
timeout: timeout,
|
||||
})
|
||||
xlog.Info("localai-proxy: ready", "upstream", raw, "upstream_model", model,
|
||||
"has_key", key != "", "realtime_pipeline", pipeline)
|
||||
return nil
|
||||
}
|
||||
|
||||
// config returns the loaded configuration, or the typed not-loaded error so
|
||||
// callers see FailedPrecondition instead of a nil dereference.
|
||||
func (p *LocalAIProxy) config() (*proxyConfig, error) {
|
||||
cfg := p.cfg.Load()
|
||||
if cfg == nil {
|
||||
return nil, grpcerrors.ModelNotLoaded(backendName)
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// model returns the model name to send upstream. The configured name wins so
|
||||
// every method targets the same upstream model; req (a model named by the
|
||||
// request itself) is only a fallback for configs that resolved no name.
|
||||
func (p *LocalAIProxy) model(req string) string {
|
||||
if cfg := p.cfg.Load(); cfg != nil && cfg.upstreamModel != "" {
|
||||
return cfg.upstreamModel
|
||||
}
|
||||
return req
|
||||
}
|
||||
|
||||
// resolveAPIKey mirrors config.ProxyConfig.ResolveAPIKey (and cloud-proxy's
|
||||
// copy). Duplicated so the backend binary does not depend on core's layout.
|
||||
func resolveAPIKey(envName, filePath string) (string, error) {
|
||||
if envName != "" {
|
||||
v := os.Getenv(envName)
|
||||
if v == "" {
|
||||
return "", fmt.Errorf("localai-proxy: api_key_env %q is unset", envName)
|
||||
}
|
||||
return v, nil
|
||||
}
|
||||
if filePath != "" {
|
||||
b, err := os.ReadFile(filePath)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("localai-proxy: read api_key_file %q: %w", filePath, err)
|
||||
}
|
||||
return strings.TrimSpace(string(b)), nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// unimplemented is the error for methods LocalAI's REST API cannot serve.
|
||||
// Failover reads gRPC Unimplemented as a capability gap and moves to the next
|
||||
// target without marking this one unhealthy.
|
||||
func unimplemented(method string) error {
|
||||
return status.Errorf(codes.Unimplemented, "localai-proxy: %s has no upstream counterpart", method)
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) AudioEncode(*pb.AudioEncodeRequest) (*pb.AudioEncodeResult, error) {
|
||||
return nil, unimplemented("AudioEncode")
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) AudioDecode(*pb.AudioDecodeRequest) (*pb.AudioDecodeResult, error) {
|
||||
return nil, unimplemented("AudioDecode")
|
||||
}
|
||||
|
||||
// AudioToAudioStream closes out because the gRPC server drains it until
|
||||
// closed; leaving it open would hang the call.
|
||||
func (p *LocalAIProxy) AudioToAudioStream(_ <-chan *pb.AudioToAudioRequest, out chan<- *pb.AudioToAudioResponse) error {
|
||||
close(out)
|
||||
return unimplemented("AudioToAudioStream")
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) TokenClassify(context.Context, *pb.TokenClassifyRequest) (*pb.TokenClassifyResponse, error) {
|
||||
return nil, unimplemented("TokenClassify")
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) ModelMetadata(*pb.ModelOptions) (*pb.ModelMetadataResponse, error) {
|
||||
return nil, unimplemented("ModelMetadata")
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) StartFineTune(*pb.FineTuneRequest) (*pb.FineTuneJobResult, error) {
|
||||
return nil, unimplemented("StartFineTune")
|
||||
}
|
||||
|
||||
// FineTuneProgress closes the channel: the gRPC server waits for it to close
|
||||
// before returning, and base.Base leaves it open.
|
||||
func (p *LocalAIProxy) FineTuneProgress(_ *pb.FineTuneProgressRequest, updates chan *pb.FineTuneProgressUpdate) error {
|
||||
close(updates)
|
||||
return unimplemented("FineTuneProgress")
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) StopFineTune(*pb.FineTuneStopRequest) error {
|
||||
return unimplemented("StopFineTune")
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) ListCheckpoints(*pb.ListCheckpointsRequest) (*pb.ListCheckpointsResponse, error) {
|
||||
return nil, unimplemented("ListCheckpoints")
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) ExportModel(*pb.ExportModelRequest) error {
|
||||
return unimplemented("ExportModel")
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) StartQuantization(*pb.QuantizationRequest) (*pb.QuantizationJobResult, error) {
|
||||
return nil, unimplemented("StartQuantization")
|
||||
}
|
||||
|
||||
// QuantizationProgress closes the channel for the same reason as
|
||||
// FineTuneProgress.
|
||||
func (p *LocalAIProxy) QuantizationProgress(_ *pb.QuantizationProgressRequest, updates chan *pb.QuantizationProgressUpdate) error {
|
||||
close(updates)
|
||||
return unimplemented("QuantizationProgress")
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) StopQuantization(*pb.QuantizationStopRequest) error {
|
||||
return unimplemented("StopQuantization")
|
||||
}
|
||||
Executable
+6
@@ -0,0 +1,6 @@
|
||||
#!/bin/bash
|
||||
set -ex
|
||||
|
||||
CURDIR=$(dirname "$(realpath "$0")")
|
||||
|
||||
exec "$CURDIR"/localai-proxy "$@"
|
||||
@@ -0,0 +1,362 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/xlog"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
// textRequest is the body for /v1/chat/completions (Messages) and
|
||||
// /v1/completions (Prompt). Zero sampling values are omitted so the upstream
|
||||
// model's own config defaults apply, as they would for a direct caller.
|
||||
type textRequest struct {
|
||||
Model string `json:"model"`
|
||||
Messages []chatMessage `json:"messages,omitempty"`
|
||||
Prompt string `json:"prompt,omitempty"`
|
||||
Stream bool `json:"stream,omitempty"`
|
||||
MaxTokens int32 `json:"max_tokens,omitempty"`
|
||||
Temperature float32 `json:"temperature,omitempty"`
|
||||
TopP float32 `json:"top_p,omitempty"`
|
||||
TopK int32 `json:"top_k,omitempty"`
|
||||
Seed int32 `json:"seed,omitempty"`
|
||||
Stop []string `json:"stop,omitempty"`
|
||||
Tools json.RawMessage `json:"tools,omitempty"`
|
||||
ToolChoice json.RawMessage `json:"tool_choice,omitempty"`
|
||||
}
|
||||
|
||||
type chatMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
Name string `json:"name,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||
ToolCalls []toolCall `json:"tool_calls,omitempty"`
|
||||
}
|
||||
|
||||
type toolCall struct {
|
||||
Index int `json:"index"`
|
||||
ID string `json:"id,omitempty"`
|
||||
Type string `json:"type,omitempty"`
|
||||
Function struct {
|
||||
Name string `json:"name,omitempty"`
|
||||
Arguments string `json:"arguments,omitempty"`
|
||||
} `json:"function"`
|
||||
}
|
||||
|
||||
// textChoice covers both endpoints and both shapes: chat replies fill Message
|
||||
// (or Delta when streaming), completions fill Text.
|
||||
type textChoice struct {
|
||||
Text string `json:"text"`
|
||||
Message choiceDelta `json:"message"`
|
||||
Delta choiceDelta `json:"delta"`
|
||||
}
|
||||
|
||||
type choiceDelta struct {
|
||||
Content string `json:"content"`
|
||||
// LocalAI names the field "reasoning"; other OpenAI-compatible servers
|
||||
// use "reasoning_content".
|
||||
Reasoning string `json:"reasoning"`
|
||||
ReasoningContent string `json:"reasoning_content"`
|
||||
ToolCalls []toolCall `json:"tool_calls"`
|
||||
}
|
||||
|
||||
type textResponse struct {
|
||||
Choices []textChoice `json:"choices"`
|
||||
Usage *struct {
|
||||
PromptTokens int32 `json:"prompt_tokens"`
|
||||
CompletionTokens int32 `json:"completion_tokens"`
|
||||
} `json:"usage"`
|
||||
}
|
||||
|
||||
// textRequest picks the chat endpoint when core sent structured messages
|
||||
// (the model uses the tokenizer template, so the upstream must template
|
||||
// too); otherwise core already rendered the prompt and completions takes it
|
||||
// verbatim.
|
||||
func (p *LocalAIProxy) textRequest(opts *pb.PredictOptions, stream bool) (string, textRequest) {
|
||||
req := textRequest{
|
||||
Model: p.model(""),
|
||||
Stream: stream,
|
||||
MaxTokens: opts.GetTokens(),
|
||||
Temperature: opts.GetTemperature(),
|
||||
TopP: opts.GetTopP(),
|
||||
TopK: opts.GetTopK(),
|
||||
Seed: opts.GetSeed(),
|
||||
Stop: opts.GetStopPrompts(),
|
||||
Tools: rawJSON(opts.GetTools()),
|
||||
ToolChoice: rawJSON(opts.GetToolChoice()),
|
||||
}
|
||||
if len(opts.GetMessages()) == 0 {
|
||||
req.Prompt = opts.GetPrompt()
|
||||
return "/v1/completions", req
|
||||
}
|
||||
for _, m := range opts.GetMessages() {
|
||||
msg := chatMessage{
|
||||
Role: m.GetRole(),
|
||||
Content: m.GetContent(),
|
||||
Name: m.GetName(),
|
||||
ToolCallID: m.GetToolCallId(),
|
||||
}
|
||||
// A previous assistant turn carries its tool calls as a JSON string.
|
||||
if tc := m.GetToolCalls(); tc != "" {
|
||||
if err := json.Unmarshal([]byte(tc), &msg.ToolCalls); err != nil {
|
||||
xlog.Debug("localai-proxy: drop malformed tool_calls on message", "error", err)
|
||||
}
|
||||
}
|
||||
req.Messages = append(req.Messages, msg)
|
||||
}
|
||||
return "/v1/chat/completions", req
|
||||
}
|
||||
|
||||
// rawJSON passes a JSON string through untouched, or omits it when it is
|
||||
// empty or invalid rather than failing the whole request.
|
||||
func rawJSON(s string) json.RawMessage {
|
||||
if s == "" || !json.Valid([]byte(s)) {
|
||||
return nil
|
||||
}
|
||||
return json.RawMessage(s)
|
||||
}
|
||||
|
||||
// replyFromChoice builds the Reply for one choice. Reasoning and tool calls
|
||||
// travel as ChatDeltas, the same shape the llama.cpp autoparser emits, so
|
||||
// core handles them without parsing the text again.
|
||||
func replyFromChoice(c textChoice, streaming bool) *pb.Reply {
|
||||
d := c.Message
|
||||
if streaming {
|
||||
d = c.Delta
|
||||
}
|
||||
content := d.Content
|
||||
if content == "" {
|
||||
content = c.Text
|
||||
}
|
||||
reasoning := d.Reasoning
|
||||
if reasoning == "" {
|
||||
reasoning = d.ReasoningContent
|
||||
}
|
||||
|
||||
reply := &pb.Reply{Message: []byte(content)}
|
||||
// Streaming chunks always carry a delta, like the autoparser's. A
|
||||
// complete reply only needs one for reasoning or tool calls; its
|
||||
// content must then be in the delta too, because core takes the
|
||||
// content from the deltas once they carry anything.
|
||||
if reasoning == "" && len(d.ToolCalls) == 0 && (!streaming || content == "") {
|
||||
return reply
|
||||
}
|
||||
delta := &pb.ChatDelta{Content: content, ReasoningContent: reasoning}
|
||||
for _, tc := range d.ToolCalls {
|
||||
delta.ToolCalls = append(delta.ToolCalls, &pb.ToolCallDelta{
|
||||
Index: int32(tc.Index),
|
||||
Id: tc.ID,
|
||||
Name: tc.Function.Name,
|
||||
Arguments: tc.Function.Arguments,
|
||||
})
|
||||
}
|
||||
reply.ChatDeltas = []*pb.ChatDelta{delta}
|
||||
return reply
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) PredictRich(opts *pb.PredictOptions) (*pb.Reply, error) {
|
||||
path, body := p.textRequest(opts, false)
|
||||
var resp textResponse
|
||||
if err := p.postJSON(context.Background(), path, body, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(resp.Choices) == 0 {
|
||||
return nil, status.Errorf(codes.Internal, "localai-proxy: upstream %s returned no choices", path)
|
||||
}
|
||||
reply := replyFromChoice(resp.Choices[0], false)
|
||||
if resp.Usage != nil {
|
||||
reply.PromptTokens = resp.Usage.PromptTokens
|
||||
reply.Tokens = resp.Usage.CompletionTokens
|
||||
}
|
||||
return reply, nil
|
||||
}
|
||||
|
||||
// PredictStreamRich sends one Reply per upstream SSE delta. It does not close
|
||||
// results: the gRPC server does, after this returns.
|
||||
func (p *LocalAIProxy) PredictStreamRich(opts *pb.PredictOptions, results chan<- *pb.Reply) error {
|
||||
path, body := p.textRequest(opts, true)
|
||||
resp, err := p.postStream(context.Background(), path, body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
// A single frame can carry a long tool-call argument chunk.
|
||||
scanner.Buffer(make([]byte, 0, 64*1024), 4<<20)
|
||||
for scanner.Scan() {
|
||||
payload, ok := strings.CutPrefix(scanner.Text(), "data:")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
payload = strings.TrimSpace(payload)
|
||||
if payload == "[DONE]" {
|
||||
return nil
|
||||
}
|
||||
if payload == "" {
|
||||
continue
|
||||
}
|
||||
var chunk textResponse
|
||||
if err := json.Unmarshal([]byte(payload), &chunk); err != nil {
|
||||
xlog.Debug("localai-proxy: skip malformed SSE frame", "path", path, "error", err)
|
||||
continue
|
||||
}
|
||||
if chunk.Usage != nil && len(chunk.Choices) == 0 {
|
||||
results <- &pb.Reply{PromptTokens: chunk.Usage.PromptTokens, Tokens: chunk.Usage.CompletionTokens}
|
||||
continue
|
||||
}
|
||||
for _, c := range chunk.Choices {
|
||||
reply := replyFromChoice(c, true)
|
||||
if len(reply.GetMessage()) == 0 && len(reply.GetChatDeltas()) == 0 {
|
||||
continue // role-only or finish frames carry nothing to emit
|
||||
}
|
||||
results <- reply
|
||||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
return transportError(path, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Predict is the legacy string path; the gRPC server prefers PredictRich.
|
||||
func (p *LocalAIProxy) Predict(opts *pb.PredictOptions) (string, error) {
|
||||
reply, err := p.PredictRich(opts)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(reply.GetMessage()), nil
|
||||
}
|
||||
|
||||
// PredictStream is the legacy string stream. Unlike PredictStreamRich it
|
||||
// owns and closes results, per the AIModel contract.
|
||||
func (p *LocalAIProxy) PredictStream(opts *pb.PredictOptions, results chan string) error {
|
||||
defer close(results)
|
||||
rich := make(chan *pb.Reply)
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
errCh <- p.PredictStreamRich(opts, rich)
|
||||
close(rich)
|
||||
}()
|
||||
for reply := range rich {
|
||||
if msg := reply.GetMessage(); len(msg) > 0 {
|
||||
results <- string(msg)
|
||||
}
|
||||
}
|
||||
return <-errCh
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) Embeddings(opts *pb.PredictOptions) ([]float32, error) {
|
||||
var resp struct {
|
||||
Data []struct {
|
||||
Embedding []float32 `json:"embedding"`
|
||||
} `json:"data"`
|
||||
}
|
||||
body := map[string]any{"model": p.model(""), "input": opts.GetEmbeddings()}
|
||||
if err := p.postJSON(context.Background(), "/v1/embeddings", body, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(resp.Data) == 0 || len(resp.Data[0].Embedding) == 0 {
|
||||
return nil, status.Error(codes.Internal, "localai-proxy: upstream /v1/embeddings returned no embedding")
|
||||
}
|
||||
return resp.Data[0].Embedding, nil
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) Rerank(ctx context.Context, in *pb.RerankRequest) (*pb.RerankResult, error) {
|
||||
body := map[string]any{
|
||||
"model": p.model(""),
|
||||
"query": in.GetQuery(),
|
||||
"documents": in.GetDocuments(),
|
||||
"top_n": in.GetTopN(),
|
||||
}
|
||||
var resp struct {
|
||||
Usage struct {
|
||||
TotalTokens int32 `json:"total_tokens"`
|
||||
PromptTokens int32 `json:"prompt_tokens"`
|
||||
} `json:"usage"`
|
||||
Results []struct {
|
||||
Index int32 `json:"index"`
|
||||
Document struct {
|
||||
Text string `json:"text"`
|
||||
} `json:"document"`
|
||||
RelevanceScore float32 `json:"relevance_score"`
|
||||
} `json:"results"`
|
||||
}
|
||||
if err := p.postJSON(ctx, "/v1/rerank", body, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := &pb.RerankResult{Usage: &pb.Usage{TotalTokens: resp.Usage.TotalTokens, PromptTokens: resp.Usage.PromptTokens}}
|
||||
for _, r := range resp.Results {
|
||||
out.Results = append(out.Results, &pb.DocumentResult{Index: r.Index, Text: r.Document.Text, RelevanceScore: r.RelevanceScore})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// TokenizeString uses the upstream model's tokenizer, which is the one that
|
||||
// matters: the upstream is where the tokens will be spent.
|
||||
func (p *LocalAIProxy) TokenizeString(opts *pb.PredictOptions) (pb.TokenizationResponse, error) {
|
||||
var resp struct {
|
||||
Tokens []int32 `json:"tokens"`
|
||||
}
|
||||
body := map[string]any{"model": p.model(""), "content": opts.GetPrompt()}
|
||||
if err := p.postJSON(context.Background(), "/v1/tokenize", body, &resp); err != nil {
|
||||
return pb.TokenizationResponse{}, err
|
||||
}
|
||||
return pb.TokenizationResponse{Length: int32(len(resp.Tokens)), Tokens: resp.Tokens}, nil
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) Detokenize(in *pb.DetokenizeRequest) (pb.DetokenizeResponse, error) {
|
||||
var resp struct {
|
||||
Content string `json:"content"`
|
||||
}
|
||||
body := map[string]any{"model": p.model(""), "tokens": in.GetTokens()}
|
||||
if err := p.postJSON(context.Background(), "/v1/detokenize", body, &resp); err != nil {
|
||||
return pb.DetokenizeResponse{}, err
|
||||
}
|
||||
return pb.DetokenizeResponse{Content: resp.Content}, nil
|
||||
}
|
||||
|
||||
// Score forwards plain candidate scoring to /api/score. That endpoint has no
|
||||
// decision-pipeline fields, so a question_type request would silently become
|
||||
// plain scoring upstream; refuse it as a capability gap instead.
|
||||
func (p *LocalAIProxy) Score(ctx context.Context, in *pb.ScoreRequest) (*pb.ScoreResponse, error) {
|
||||
if in.GetQuestionType() != "" {
|
||||
return nil, unimplemented("Score with question_type")
|
||||
}
|
||||
body := map[string]any{
|
||||
"model": p.model(""),
|
||||
"prompt": in.GetPrompt(),
|
||||
"candidates": in.GetCandidates(),
|
||||
"include_token_logprobs": in.GetIncludeTokenLogprobs(),
|
||||
"length_normalize": in.GetLengthNormalize(),
|
||||
}
|
||||
var resp struct {
|
||||
Candidates []struct {
|
||||
LogProb float64 `json:"log_prob"`
|
||||
LengthNormalizedLogProb float64 `json:"length_normalized_log_prob"`
|
||||
NumTokens int32 `json:"num_tokens"`
|
||||
Tokens []struct {
|
||||
Token string `json:"token"`
|
||||
LogProb float64 `json:"log_prob"`
|
||||
} `json:"tokens"`
|
||||
} `json:"candidates"`
|
||||
}
|
||||
if err := p.postJSON(ctx, "/api/score", body, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := &pb.ScoreResponse{}
|
||||
for _, c := range resp.Candidates {
|
||||
cs := &pb.CandidateScore{LogProb: c.LogProb, LengthNormalizedLogProb: c.LengthNormalizedLogProb, NumTokens: c.NumTokens}
|
||||
for _, t := range c.Tokens {
|
||||
cs.Tokens = append(cs.Tokens, &pb.TokenLogProb{Token: t.Token, LogProb: t.LogProb})
|
||||
}
|
||||
out.Candidates = append(out.Candidates, cs)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
@@ -0,0 +1,483 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
grpc "github.com/mudler/LocalAI/pkg/grpc"
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
func codeOf(err error) codes.Code {
|
||||
st, ok := status.FromError(err)
|
||||
if !ok {
|
||||
return codes.Unknown
|
||||
}
|
||||
return st.Code()
|
||||
}
|
||||
|
||||
var _ = Describe("localai-proxy", func() {
|
||||
var up *fakeUpstream
|
||||
|
||||
BeforeEach(func() {
|
||||
up = newFakeUpstream()
|
||||
DeferCleanup(up.Close)
|
||||
})
|
||||
|
||||
Describe("Load", func() {
|
||||
It("refuses a model without proxy options", func() {
|
||||
err := NewLocalAIProxy().Load(&pb.ModelOptions{Model: "m"})
|
||||
Expect(err).To(MatchError(ContainSubstring("proxy")))
|
||||
})
|
||||
|
||||
It("refuses a missing or invalid upstream_url", func() {
|
||||
p := NewLocalAIProxy()
|
||||
Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{}})).NotTo(Succeed())
|
||||
Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{UpstreamUrl: "not a url"}})).NotTo(Succeed())
|
||||
Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{UpstreamUrl: "ftp://host"}})).NotTo(Succeed())
|
||||
})
|
||||
|
||||
It("refuses an api_key_env that is unset", func() {
|
||||
err := NewLocalAIProxy().Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{
|
||||
UpstreamUrl: up.URL, ApiKeyEnv: "LOCALAI_PROXY_TEST_UNSET_KEY",
|
||||
}})
|
||||
Expect(err).To(MatchError(ContainSubstring("LOCALAI_PROXY_TEST_UNSET_KEY")))
|
||||
})
|
||||
|
||||
It("parses realtime_pipeline, strips the trailing slash and keeps the timeout", func() {
|
||||
p := loadProxy(up, func(o *pb.ModelOptions) {
|
||||
o.Options = []string{"other:1", "realtime_pipeline:my-pipe"}
|
||||
o.Proxy.RequestTimeoutSeconds = 7
|
||||
})
|
||||
cfg := p.cfg.Load()
|
||||
Expect(cfg.realtimePipeline).To(Equal("my-pipe"))
|
||||
Expect(cfg.base).To(Equal(up.URL))
|
||||
Expect(cfg.timeout).To(Equal(7 * time.Second))
|
||||
})
|
||||
|
||||
It("falls back to the model name when upstream_model is unset", func() {
|
||||
p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.UpstreamModel = "" })
|
||||
Expect(p.model("")).To(Equal("local-name"))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("PredictRich", func() {
|
||||
It("sends messages to /v1/chat/completions with the upstream model and key", func() {
|
||||
GinkgoT().Setenv("LOCALAI_PROXY_TEST_KEY", "sk-test")
|
||||
p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.ApiKeyEnv = "LOCALAI_PROXY_TEST_KEY" })
|
||||
up.replyJSON("/v1/chat/completions", map[string]any{
|
||||
"choices": []any{map[string]any{"message": map[string]any{"role": "assistant", "content": "hello back"}}},
|
||||
"usage": map[string]any{"prompt_tokens": 3, "completion_tokens": 2},
|
||||
})
|
||||
|
||||
reply, err := p.PredictRich(&pb.PredictOptions{
|
||||
Messages: []*pb.Message{{Role: "user", Content: "hello"}},
|
||||
Tokens: 32,
|
||||
Temperature: 0.5,
|
||||
TopK: 40,
|
||||
StopPrompts: []string{"</s>"},
|
||||
Seed: 9,
|
||||
})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(string(reply.GetMessage())).To(Equal("hello back"))
|
||||
Expect(reply.GetPromptTokens()).To(Equal(int32(3)))
|
||||
Expect(reply.GetTokens()).To(Equal(int32(2)))
|
||||
|
||||
req := up.last()
|
||||
Expect(req.Method).To(Equal(http.MethodPost))
|
||||
Expect(req.Path).To(Equal("/v1/chat/completions"))
|
||||
Expect(req.Auth).To(Equal("Bearer sk-test"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("max_tokens", BeNumerically("==", 32)))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("temperature", BeNumerically("==", 0.5)))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("top_k", BeNumerically("==", 40)))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("seed", BeNumerically("==", 9)))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("stop", ConsistOf("</s>")))
|
||||
Expect(req.JSON).NotTo(HaveKey("stream"))
|
||||
Expect(req.JSON["messages"]).To(ConsistOf(HaveKeyWithValue("content", "hello")))
|
||||
})
|
||||
|
||||
It("returns upstream tool calls as chat deltas", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/chat/completions", map[string]any{
|
||||
"choices": []any{map[string]any{"message": map[string]any{
|
||||
"role": "assistant",
|
||||
"tool_calls": []any{map[string]any{
|
||||
"id": "call_1", "type": "function",
|
||||
"function": map[string]any{"name": "get_weather", "arguments": `{"city":"Rome"}`},
|
||||
}},
|
||||
}}},
|
||||
})
|
||||
|
||||
reply, err := p.PredictRich(&pb.PredictOptions{
|
||||
Messages: []*pb.Message{{Role: "user", Content: "weather?"}},
|
||||
Tools: `[{"type":"function","function":{"name":"get_weather"}}]`,
|
||||
})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(reply.GetChatDeltas()).To(HaveLen(1))
|
||||
tc := reply.GetChatDeltas()[0].GetToolCalls()
|
||||
Expect(tc).To(HaveLen(1))
|
||||
Expect(tc[0].GetName()).To(Equal("get_weather"))
|
||||
Expect(tc[0].GetArguments()).To(Equal(`{"city":"Rome"}`))
|
||||
Expect(up.last().JSON).To(HaveKey("tools"))
|
||||
})
|
||||
|
||||
It("sends a bare prompt to /v1/completions", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/completions", map[string]any{
|
||||
"choices": []any{map[string]any{"text": "completed"}},
|
||||
})
|
||||
|
||||
reply, err := p.PredictRich(&pb.PredictOptions{Prompt: "once upon"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(string(reply.GetMessage())).To(Equal("completed"))
|
||||
req := up.last()
|
||||
Expect(req.Path).To(Equal("/v1/completions"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("prompt", "once upon"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
|
||||
})
|
||||
|
||||
It("maps a 5xx upstream to Unavailable with the body in the message", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/chat/completions", scriptedResponse{Status: http.StatusServiceUnavailable, Body: "backend is down"})
|
||||
|
||||
_, err := p.PredictRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "x"}}})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unavailable))
|
||||
Expect(err.Error()).To(ContainSubstring("backend is down"))
|
||||
})
|
||||
|
||||
It("maps a 4xx upstream to InvalidArgument", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/chat/completions", scriptedResponse{Status: http.StatusBadRequest, Body: "bad request"})
|
||||
|
||||
_, err := p.PredictRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "x"}}})
|
||||
Expect(codeOf(err)).To(Equal(codes.InvalidArgument))
|
||||
Expect(err.Error()).To(ContainSubstring("bad request"))
|
||||
})
|
||||
|
||||
It("truncates a long upstream error body", func() {
|
||||
p := loadProxy(up, nil)
|
||||
long := make([]byte, 2000)
|
||||
for i := range long {
|
||||
long[i] = 'a'
|
||||
}
|
||||
up.script("/v1/completions", scriptedResponse{Status: http.StatusInternalServerError, Body: string(long)})
|
||||
|
||||
_, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unavailable))
|
||||
Expect(len(status.Convert(err).Message())).To(BeNumerically("<", 700))
|
||||
})
|
||||
|
||||
It("maps an unreachable upstream to Unavailable", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.Close()
|
||||
|
||||
_, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unavailable))
|
||||
})
|
||||
|
||||
It("reports an unloaded proxy as FailedPrecondition", func() {
|
||||
_, err := NewLocalAIProxy().PredictRich(&pb.PredictOptions{Prompt: "x"})
|
||||
Expect(codeOf(err)).To(Equal(codes.FailedPrecondition))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("PredictStreamRich", func() {
|
||||
It("streams SSE deltas in order and leaves the channel open", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/chat/completions", scriptedResponse{SSE: []string{
|
||||
sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"role": "assistant"}}}}),
|
||||
sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "Hel"}}}}),
|
||||
sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "lo"}}}}),
|
||||
"[DONE]",
|
||||
}})
|
||||
|
||||
results := make(chan *pb.Reply, 10)
|
||||
err := p.PredictStreamRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, results)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
var got []string
|
||||
for len(results) > 0 {
|
||||
got = append(got, string((<-results).GetMessage()))
|
||||
}
|
||||
Expect(got).To(Equal([]string{"Hel", "lo"}))
|
||||
// The gRPC server closes the channel; closing it here must not panic.
|
||||
close(results)
|
||||
Expect(up.last().JSON).To(HaveKeyWithValue("stream", true))
|
||||
})
|
||||
|
||||
It("streams /v1/completions text for a bare prompt", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/completions", scriptedResponse{SSE: []string{
|
||||
sseJSON(map[string]any{"choices": []any{map[string]any{"text": "a"}}}),
|
||||
sseJSON(map[string]any{"choices": []any{map[string]any{"text": "b"}}}),
|
||||
"[DONE]",
|
||||
}})
|
||||
|
||||
results := make(chan *pb.Reply, 10)
|
||||
Expect(p.PredictStreamRich(&pb.PredictOptions{Prompt: "go"}, results)).To(Succeed())
|
||||
Expect(results).To(HaveLen(2))
|
||||
Expect(string((<-results).GetMessage())).To(Equal("a"))
|
||||
Expect(string((<-results).GetMessage())).To(Equal("b"))
|
||||
})
|
||||
|
||||
It("maps a failing upstream to a gRPC code", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/completions", scriptedResponse{Status: http.StatusBadGateway, Body: "gateway"})
|
||||
|
||||
err := p.PredictStreamRich(&pb.PredictOptions{Prompt: "go"}, make(chan *pb.Reply, 1))
|
||||
Expect(codeOf(err)).To(Equal(codes.Unavailable))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("legacy Predict and PredictStream", func() {
|
||||
It("wrap the rich variants", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/completions", map[string]any{"choices": []any{map[string]any{"text": "plain"}}})
|
||||
out, err := p.Predict(&pb.PredictOptions{Prompt: "x"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(out).To(Equal("plain"))
|
||||
|
||||
up.script("/v1/completions", scriptedResponse{SSE: []string{
|
||||
sseJSON(map[string]any{"choices": []any{map[string]any{"text": "s1"}}}),
|
||||
"[DONE]",
|
||||
}})
|
||||
results := make(chan string, 10)
|
||||
Expect(p.PredictStream(&pb.PredictOptions{Prompt: "x"}, results)).To(Succeed())
|
||||
var got []string
|
||||
for s := range results { // PredictStream closes the channel
|
||||
got = append(got, s)
|
||||
}
|
||||
Expect(got).To(Equal([]string{"s1"}))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("Embeddings", func() {
|
||||
It("posts the input to /v1/embeddings and returns the first vector", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/embeddings", map[string]any{
|
||||
"data": []any{map[string]any{"embedding": []float32{0.1, 0.2, 0.3}}},
|
||||
})
|
||||
|
||||
vec, err := p.Embeddings(&pb.PredictOptions{Embeddings: "embed me"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(vec).To(Equal([]float32{0.1, 0.2, 0.3}))
|
||||
req := up.last()
|
||||
Expect(req.Path).To(Equal("/v1/embeddings"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("input", "embed me"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
|
||||
})
|
||||
|
||||
It("fails when the upstream returns no vector", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/embeddings", map[string]any{"data": []any{}})
|
||||
_, err := p.Embeddings(&pb.PredictOptions{Embeddings: "x"})
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
})
|
||||
|
||||
Describe("Rerank", func() {
|
||||
It("posts to /v1/rerank and maps the results", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/rerank", map[string]any{
|
||||
"model": "remote-model",
|
||||
"usage": map[string]any{"total_tokens": 12, "prompt_tokens": 10},
|
||||
"results": []any{
|
||||
map[string]any{"index": 1, "document": map[string]any{"text": "b"}, "relevance_score": 0.9},
|
||||
map[string]any{"index": 0, "document": map[string]any{"text": "a"}, "relevance_score": 0.1},
|
||||
},
|
||||
})
|
||||
|
||||
res, err := p.Rerank(context.Background(), &pb.RerankRequest{Query: "q", Documents: []string{"a", "b"}, TopN: 2})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(res.GetUsage().GetTotalTokens()).To(Equal(int32(12)))
|
||||
Expect(res.GetUsage().GetPromptTokens()).To(Equal(int32(10)))
|
||||
Expect(res.GetResults()).To(HaveLen(2))
|
||||
Expect(res.GetResults()[0].GetIndex()).To(Equal(int32(1)))
|
||||
Expect(res.GetResults()[0].GetText()).To(Equal("b"))
|
||||
Expect(res.GetResults()[0].GetRelevanceScore()).To(BeNumerically("~", 0.9, 1e-6))
|
||||
|
||||
req := up.last()
|
||||
Expect(req.Path).To(Equal("/v1/rerank"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("query", "q"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("top_n", BeNumerically("==", 2)))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("documents", ConsistOf("a", "b")))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("TokenizeString and Detokenize", func() {
|
||||
It("posts the prompt to /v1/tokenize", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/tokenize", map[string]any{"tokens": []int32{5, 6, 7}})
|
||||
|
||||
res, err := p.TokenizeString(&pb.PredictOptions{Prompt: "abc"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(res.GetTokens()).To(Equal([]int32{5, 6, 7}))
|
||||
Expect(res.GetLength()).To(Equal(int32(3)))
|
||||
req := up.last()
|
||||
Expect(req.Path).To(Equal("/v1/tokenize"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("content", "abc"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
|
||||
})
|
||||
|
||||
It("posts tokens to /v1/detokenize", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/detokenize", map[string]any{"content": "abc"})
|
||||
|
||||
res, err := p.Detokenize(&pb.DetokenizeRequest{Tokens: []int32{5, 6}})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(res.GetContent()).To(Equal("abc"))
|
||||
Expect(up.last().JSON).To(HaveKeyWithValue("tokens", ConsistOf(BeNumerically("==", 5), BeNumerically("==", 6))))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("Score", func() {
|
||||
It("posts to /api/score and maps the candidates", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/api/score", map[string]any{
|
||||
"model": "remote-model",
|
||||
"candidates": []any{map[string]any{
|
||||
"log_prob": -1.5, "length_normalized_log_prob": -0.75, "num_tokens": 2,
|
||||
"tokens": []any{map[string]any{"token": "yes", "log_prob": -1.5}},
|
||||
}},
|
||||
})
|
||||
|
||||
res, err := p.Score(context.Background(), &pb.ScoreRequest{
|
||||
Prompt: "p", Candidates: []string{"yes"}, IncludeTokenLogprobs: true, LengthNormalize: true,
|
||||
})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(res.GetCandidates()).To(HaveLen(1))
|
||||
c := res.GetCandidates()[0]
|
||||
Expect(c.GetLogProb()).To(Equal(-1.5))
|
||||
Expect(c.GetLengthNormalizedLogProb()).To(Equal(-0.75))
|
||||
Expect(c.GetNumTokens()).To(Equal(int32(2)))
|
||||
Expect(c.GetTokens()).To(HaveLen(1))
|
||||
Expect(c.GetTokens()[0].GetToken()).To(Equal("yes"))
|
||||
|
||||
req := up.last()
|
||||
Expect(req.Path).To(Equal("/api/score"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("prompt", "p"))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("include_token_logprobs", true))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("length_normalize", true))
|
||||
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
|
||||
})
|
||||
|
||||
It("refuses decision-pipeline requests it cannot forward", func() {
|
||||
p := loadProxy(up, nil)
|
||||
_, err := p.Score(context.Background(), &pb.ScoreRequest{Prompt: "{}", QuestionType: "systemone"})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(up.recorded()).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
Describe("helpers", func() {
|
||||
It("postMultipart sends fields and the file with auth", func() {
|
||||
GinkgoT().Setenv("LOCALAI_PROXY_TEST_KEY", "sk-mp")
|
||||
p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.ApiKeyEnv = "LOCALAI_PROXY_TEST_KEY" })
|
||||
up.replyJSON("/v1/audio/transcriptions", map[string]any{"text": "ok"})
|
||||
|
||||
path := filepath.Join(GinkgoT().TempDir(), "a.wav")
|
||||
Expect(os.WriteFile(path, []byte("RIFFDATA"), 0o600)).To(Succeed())
|
||||
|
||||
var out struct {
|
||||
Text string `json:"text"`
|
||||
}
|
||||
err := p.postMultipart(context.Background(), "/v1/audio/transcriptions",
|
||||
map[string]string{"model": "remote-model", "language": "it"}, "file", path, &out)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(out.Text).To(Equal("ok"))
|
||||
req := up.last()
|
||||
Expect(req.Auth).To(Equal("Bearer sk-mp"))
|
||||
Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "language": "it"}))
|
||||
Expect(req.Files).To(HaveKeyWithValue("file", "RIFFDATA"))
|
||||
})
|
||||
|
||||
It("postMultipart reports a missing local file without calling upstream", func() {
|
||||
p := loadProxy(up, nil)
|
||||
err := p.postMultipart(context.Background(), "/v1/audio/transcriptions", nil, "file", "/nonexistent/a.wav", nil)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(up.recorded()).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("applies request_timeout_seconds to non-streaming calls", func() {
|
||||
slow := make(chan struct{})
|
||||
hang := newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { <-slow })
|
||||
// Cleanups run last-in first-out: release the handler before
|
||||
// Close, which waits for in-flight requests.
|
||||
DeferCleanup(hang.Close)
|
||||
DeferCleanup(func() { close(slow) })
|
||||
|
||||
p := NewLocalAIProxy()
|
||||
Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{
|
||||
UpstreamUrl: hang.URL, UpstreamModel: "m", RequestTimeoutSeconds: 1,
|
||||
}})).To(Succeed())
|
||||
_, err := p.Embeddings(&pb.PredictOptions{Embeddings: "x"})
|
||||
Expect(codeOf(err)).To(Equal(codes.DeadlineExceeded))
|
||||
})
|
||||
|
||||
It("postStream returns the open response for a 2xx", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "WAVBYTES"})
|
||||
resp, err := p.postStream(context.Background(), "/tts", map[string]any{"input": "hi"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
Expect(resp.Header.Get("Content-Type")).To(Equal("audio/wav"))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("through the gRPC server", func() {
|
||||
It("dispatches Rerank and keeps the Unimplemented code end to end", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/rerank", map[string]any{"results": []any{
|
||||
map[string]any{"index": 0, "document": map[string]any{"text": "a"}, "relevance_score": 0.5},
|
||||
}})
|
||||
addr := "test://localai-proxy-grpc"
|
||||
grpc.Provide(addr, p)
|
||||
client := grpc.NewClient(addr, true, nil, false)
|
||||
|
||||
res, err := client.Rerank(context.Background(), &pb.RerankRequest{Query: "q", Documents: []string{"a"}})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(res.GetResults()).To(HaveLen(1))
|
||||
|
||||
_, err = client.AudioEncode(context.Background(), &pb.AudioEncodeRequest{})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(status.Convert(err).Message()).To(Equal("localai-proxy: AudioEncode has no upstream counterpart"))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("methods with no upstream counterpart", func() {
|
||||
It("return Unimplemented with the exact message", func() {
|
||||
p := loadProxy(up, nil)
|
||||
_, err := p.AudioEncode(&pb.AudioEncodeRequest{})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(status.Convert(err).Message()).To(Equal("localai-proxy: AudioEncode has no upstream counterpart"))
|
||||
|
||||
_, err = p.TokenClassify(context.Background(), &pb.TokenClassifyRequest{})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
|
||||
_, err = p.ModelMetadata(&pb.ModelOptions{})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
|
||||
_, err = p.StartFineTune(&pb.FineTuneRequest{})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
|
||||
})
|
||||
|
||||
It("close the output channel of streaming stubs so the server does not hang", func() {
|
||||
p := loadProxy(up, nil)
|
||||
updates := make(chan *pb.FineTuneProgressUpdate)
|
||||
Expect(codeOf(p.FineTuneProgress(&pb.FineTuneProgressRequest{}, updates))).To(Equal(codes.Unimplemented))
|
||||
Eventually(updates).Should(BeClosed())
|
||||
|
||||
out := make(chan *pb.AudioToAudioResponse)
|
||||
Expect(codeOf(p.AudioToAudioStream(make(chan *pb.AudioToAudioRequest), out))).To(Equal(codes.Unimplemented))
|
||||
Eventually(out).Should(BeClosed())
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -2038,6 +2038,21 @@
|
||||
capabilities:
|
||||
default: "cpu-cloud-proxy"
|
||||
metal: "metal-cloud-proxy"
|
||||
- &localai-proxy
|
||||
name: "localai-proxy"
|
||||
alias: "localai-proxy"
|
||||
urls:
|
||||
- https://github.com/mudler/LocalAI/tree/master/backend/go/localai-proxy
|
||||
description: |
|
||||
Serve a model from another LocalAI instance: text, embeddings, rerank, audio, image and video requests are forwarded to its REST API.
|
||||
tags:
|
||||
- text-to-text
|
||||
- proxy
|
||||
- CPU
|
||||
license: MIT
|
||||
capabilities:
|
||||
default: "cpu-localai-proxy"
|
||||
metal: "metal-localai-proxy"
|
||||
- &valkey-store
|
||||
name: "valkey-store"
|
||||
urls:
|
||||
@@ -2623,6 +2638,31 @@
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-metal-darwin-arm64-cloud-proxy"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-metal-darwin-arm64-cloud-proxy
|
||||
- !!merge <<: *localai-proxy
|
||||
name: "cpu-localai-proxy"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-cpu-localai-proxy"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-cpu-localai-proxy
|
||||
- !!merge <<: *localai-proxy
|
||||
name: "cpu-localai-proxy-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-cpu-localai-proxy"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-cpu-localai-proxy
|
||||
- !!merge <<: *localai-proxy
|
||||
name: "localai-proxy-development"
|
||||
capabilities:
|
||||
default: "cpu-localai-proxy-development"
|
||||
metal: "metal-localai-proxy-development"
|
||||
- !!merge <<: *localai-proxy
|
||||
name: "metal-localai-proxy"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-metal-darwin-arm64-localai-proxy"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-metal-darwin-arm64-localai-proxy
|
||||
- !!merge <<: *localai-proxy
|
||||
name: "metal-localai-proxy-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-metal-darwin-arm64-localai-proxy"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-metal-darwin-arm64-localai-proxy
|
||||
- !!merge <<: *valkey-store
|
||||
name: "cpu-valkey-store"
|
||||
alias: "valkey-store"
|
||||
|
||||
Reference in new issue
Block a user