Files
LocalAI/pkg/grpc/server.go
Richard Palethorpe 9058a2bb46 feat: Add 3d generation UI/API and trellis2cpp backend (#10979)
* feat(3d): add Generate3D RPC, FLAG_3D capability, and /v1/3d/generations endpoint

Adds the plumbing for image-conditioned 3D asset generation (binary
glTF / GLB output), modeled on the video generation path:

- backend.proto: Generate3D RPC + Generate3DRequest (staged image src,
  glb dst, seed/step/cfg_scale/texture_steps, quality and background
  enums, params map for backend-specific extras)
- pkg/grpc: thread Generate3D through client, server, embed, base and
  the backend interfaces; connection-evicting and distributed-node
  wrappers (in-flight tracking + file staging) included
- core/config: FLAG_3D usecase (guessed only for the trellis2cpp
  backend), '3d' canonical usecase string mapped to the Generate3D
  method, and a '3d' output modality
- REST: POST /v1/3d/generations (+ unversioned alias) returning
  OpenAIResponse with a /generated-3d URL or b64_json; conditioning
  image accepted as URL, base64, or data URI; quality/background
  validated at the edge; .glb served as model/gltf-binary
- auth: '3d' route feature (default ON); /api/instructions entry

Assisted-by: Claude:claude-fable-5 [Claude Code]
Signed-off-by: Richard Palethorpe <io@richiejp.com>

* feat(trellis2cpp): add the trellis2.cpp image-to-3D backend

Wraps localai-org/trellis2cpp (C++/GGML port of Microsoft TRELLIS.2,
pbr-textures branch) as a Go+purego backend, following the
stablediffusion-ggml pattern:

- backend/go/trellis2cpp: purego bindings to the flat C ABI (v9,
  asserted at startup), eager pipeline load with model-set validation
  (refuses non-trellis GGUFs; degrades coarse/geometry-only/textured
  exactly like the upstream demo), Generate3D via t2_generate +
  t2_bake_glb writing a binary glTF to dst. Weight-free unit tests
  cover resolution/validation/param mapping — CI never downloads the
  multi-GB GGUF set or runs inference.
- CPU SIMD variants build into per-variant directories (the shared
  libggml sonames collide across variants, unlike sd-ggml's flat
  renamed-.so scheme); run.sh picks one via /proc/cpuinfo.
- CI wiring: backend-matrix entries (cpu, cuda12/13, vulkan
  amd64+arm64, l4t, l4t-cuda13, darwin metal), index.yaml meta +
  latest/master image entries, bump_deps tracking of the pbr-textures
  branch, changed-backends.js mapping, top-level Makefile targets.
- Importer: auto-detects trellis GGUF repos/URIs (registered before
  llama-cpp so the .gguf match isn't stolen) and expands any trellis
  URI to the full 10-file component set spanning the three LocalAI-io
  HF repos.
- Gallery: trellis2-4b (full PBR + 1024 cascade) and
  trellis2-4b-geometry (512 untextured) with verified sha256s.

Assisted-by: Claude:claude-fable-5 [Claude Code]
Signed-off-by: Richard Palethorpe <io@richiejp.com>

* feat(ui): 3D generation page with native GLB viewer and IndexedDB history

Adds a Studio tab + /app/3d page for the new image-to-3D endpoint:

- GlbViewer ports the trellis2cpp demo's dependency-free WebGL2
  renderer (quaternion trackball, metallic-roughness PBR, ACES,
  hidden-line wireframe with a bounded index budget) and pairs it with
  a minimal GLB parser for the two forms t2_bake_glb emits — dense
  vertex-PBR (linear COLOR_0 + _METALLIC_ROUGHNESS, uploaded as
  normalized integers) and the opt-in UV-atlas textured form. Parsing
  happens before any GL so stats and errors render without WebGL2.
- use3DHistory stores past generations (params, input thumbnail, and
  the GLB blob itself) in IndexedDB with keep-newest-20 eviction —
  GLBs are multi-MB binaries localStorage can't hold — and the page
  offers a download button for the active GLB.
- Wiring: CAP_3D capability constant (FLAG_3D — the exact string
  /api/models/capabilities serves), threeDApi, router entries, Studio
  tab, vite dev proxy, en locale keys.
- e2e: render-smoke entry plus a focused spec that feeds a real
  one-triangle vertex-PBR GLB through the parser/viewer and exercises
  IndexedDB persistence, selection, deletion, and API errors.

Assisted-by: Claude:claude-fable-5 [Claude Code]
Signed-off-by: Richard Palethorpe <io@richiejp.com>

* fix(3d): address API correctness and UX issues

Keep 3D generation on the LocalAI-specific /3d/generations route and ensure authentication and permissions cover it.

Propagate distributed transfer failures, publish a portable ARM64 backend image, honor importer overrides, and align discovery, upload validation, and touch controls.

Assisted-by: Codex:gpt-5
Signed-off-by: Richard Palethorpe <io@richiejp.com>

* feat(3d): add previewable print remeshing

Add a single-detail CGAL Alpha Wrap workflow for existing Trellis GLBs, including PBR reprojection, API documentation, tracing, and an in-browser preview before download.

Allow the remesh route to enforce its 512 MiB upload cap independently of the smaller global default so generated high-resolution meshes can be processed.

Assisted-by: Codex:gpt-5
Signed-off-by: Richard Palethorpe <io@richiejp.com>

* build(trellis2cpp): centralize remesh dependency pins

Assisted-by: Codex:GPT-5 [apply_patch] [exec_command]
Signed-off-by: Richard Palethorpe <io@richiejp.com>

* fix(kokoros): implement Generate3D stub for new proto RPC

The Generate3D RPC added to backend.proto for the trellis2cpp backend
made tonic's generated Backend trait require generate3_d, breaking the
kokoros-grpc build. Return unimplemented like the other unsupported
modalities.

Assisted-by: Claude Code:claude-fable-5
Signed-off-by: Richard Palethorpe <io@richiejp.com>

---------

Signed-off-by: Richard Palethorpe <io@richiejp.com>
Co-authored-by: localai-org-maint-bot <bot-opensource@localaisrl.com>
2026-07-29 16:15:04 +02:00

1106 lines
28 KiB
Go

package grpc
import (
"context"
"crypto/subtle"
"errors"
"fmt"
"io"
"log"
"net"
"os"
"strings"
"sync"
"github.com/mudler/LocalAI/pkg/grpc/grpcerrors"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/status"
)
// A GRPC Server that allows to run LLM inference.
// It is used by the LLMServices to expose the LLM functionalities that are called by the client.
// The GRPC Service is general, trying to encompass all the possible LLM options models.
// It depends on the real implementer then what can be done or not.
//
// The server is implemented as a GRPC service, with the following methods:
// - Predict: to run the inference with options
// - PredictStream: to run the inference with options and stream the results
// server is used to implement helloworld.GreeterServer.
type server struct {
pb.UnimplementedBackendServer
llm AIModel
// identityMu guards loadedIdentity: LoadModel writes it, the inference
// RPCs read it, and nothing serialises those against each other
// (llm.Locking() is the model's own lock, and it is optional).
identityMu sync.RWMutex
loadedIdentity string
}
// identifiedRequest is satisfied by every request message that carries a
// ModelIdentity field. protoc-gen-go generates GetModelIdentity with a
// nil-receiver guard, so a typed-nil request is safe to pass here.
type identifiedRequest interface {
GetModelIdentity() string
}
// checkModelIdentity reports an error when the request names a model other
// than the one this process loaded. It is the point-of-use half of the fix for
// #10952: in distributed mode a worker can recycle a stopped backend's gRPC
// port for another model's backend, and the controller's liveness-only health
// probe cannot tell a stale cached route from a valid one, so the backend has
// to be the one to catch it.
//
// Either side being empty means "skip the check". The request side is empty
// for a controller that predates the field and for internally synthesized
// requests; the loaded side is empty when such a controller performed the
// load. Neither can judge the other, and a false rejection is far worse than
// the miss.
//
// One guard covers every modality because they all share one exposure: the
// route is cached by address, the address can be recycled, and only the
// backend can tell. Which RPCs call it is the whole enforcement surface — see
// pkg/grpc/model_identity_modalities_test.go, which drives all of them.
func (s *server) checkModelIdentity(in identifiedRequest) error {
if in == nil || in.GetModelIdentity() == "" {
return nil
}
s.identityMu.RLock()
loaded := s.loadedIdentity
s.identityMu.RUnlock()
if loaded == "" || loaded == in.GetModelIdentity() {
return nil
}
return grpcerrors.ModelMismatch("grpc-server", loaded, in.GetModelIdentity())
}
func (s *server) Health(ctx context.Context, in *pb.HealthMessage) (*pb.Reply, error) {
return newReply("OK"), nil
}
func (s *server) Embedding(ctx context.Context, in *pb.PredictOptions) (*pb.EmbeddingResult, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
embeds, err := s.llm.Embeddings(in)
if err != nil {
return nil, err
}
return &pb.EmbeddingResult{Embeddings: embeds}, nil
}
func (s *server) LoadModel(ctx context.Context, in *pb.ModelOptions) (*pb.Result, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.Load(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error loading model: %s", err.Error()), Success: false}, err
}
// Record what we loaded so the PredictOptions RPCs can reject requests
// meant for a different model. Only on success: a failed load leaves no
// model, which IsModelNotLoaded already covers.
s.identityMu.Lock()
s.loadedIdentity = in.Model
s.identityMu.Unlock()
return &pb.Result{Message: "Loading succeeded", Success: true}, nil
}
func (s *server) Predict(ctx context.Context, in *pb.PredictOptions) (*pb.Reply, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
if rich, ok := s.llm.(AIModelRich); ok {
return rich.PredictRich(in)
}
result, err := s.llm.Predict(in)
return newReply(result), err
}
func (s *server) GenerateImage(ctx context.Context, in *pb.GenerateImageRequest) (*pb.Result, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.GenerateImage(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error generating image: %s", err.Error()), Success: false}, err
}
return &pb.Result{Message: "Image generated", Success: true}, nil
}
func (s *server) GenerateVideo(ctx context.Context, in *pb.GenerateVideoRequest) (*pb.Result, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.GenerateVideo(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error generating video: %s", err.Error()), Success: false}, err
}
return &pb.Result{Message: "Video generated", Success: true}, nil
}
func (s *server) Generate3D(ctx context.Context, in *pb.Generate3DRequest) (*pb.Result, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.Generate3D(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error generating 3D asset: %s", err.Error()), Success: false}, err
}
return &pb.Result{Message: "3D asset generated", Success: true}, nil
}
func (s *server) TTS(ctx context.Context, in *pb.TTSRequest) (*pb.Result, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.TTS(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error generating audio: %s", err.Error()), Success: false}, err
}
return &pb.Result{Message: "TTS audio generated", Success: true}, nil
}
func (s *server) TTSStream(in *pb.TTSRequest, stream pb.Backend_TTSStreamServer) error {
if err := s.checkModelIdentity(in); err != nil {
return err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
audioChan := make(chan []byte)
done := make(chan bool)
go func() {
for audioChunk := range audioChan {
stream.Send(&pb.Reply{Audio: audioChunk})
}
done <- true
}()
err := s.llm.TTSStream(in, audioChan)
<-done
return err
}
func (s *server) SoundGeneration(ctx context.Context, in *pb.SoundGenerationRequest) (*pb.Result, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.SoundGeneration(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error generating audio: %s", err.Error()), Success: false}, err
}
return &pb.Result{Message: "Sound Generation audio generated", Success: true}, nil
}
func (s *server) Detect(ctx context.Context, in *pb.DetectOptions) (*pb.DetectResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.Detect(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) Depth(ctx context.Context, in *pb.DepthRequest) (*pb.DepthResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.Depth(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) FaceVerify(ctx context.Context, in *pb.FaceVerifyRequest) (*pb.FaceVerifyResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.FaceVerify(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) FaceAnalyze(ctx context.Context, in *pb.FaceAnalyzeRequest) (*pb.FaceAnalyzeResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.FaceAnalyze(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) VoiceVerify(ctx context.Context, in *pb.VoiceVerifyRequest) (*pb.VoiceVerifyResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.VoiceVerify(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) VoiceAnalyze(ctx context.Context, in *pb.VoiceAnalyzeRequest) (*pb.VoiceAnalyzeResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.VoiceAnalyze(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) VoiceEmbed(ctx context.Context, in *pb.VoiceEmbedRequest) (*pb.VoiceEmbedResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.VoiceEmbed(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) AudioTranscription(ctx context.Context, in *pb.TranscriptRequest) (*pb.TranscriptResult, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
result, err := s.llm.AudioTranscription(ctx, in)
if err != nil {
return nil, err
}
tresult := &pb.TranscriptResult{}
for _, s := range result.Segments {
tks := []int32{}
for _, t := range s.Tokens {
tks = append(tks, int32(t))
}
words := make([]*pb.TranscriptWord, 0, len(s.Words))
for _, w := range s.Words {
words = append(words, &pb.TranscriptWord{
Start: int64(w.Start),
End: int64(w.End),
Text: w.Text,
})
}
tresult.Segments = append(tresult.Segments,
&pb.TranscriptSegment{
Text: s.Text,
Id: int32(s.Id),
Start: int64(s.Start),
End: int64(s.End),
Tokens: tks,
Speaker: s.Speaker,
Words: words,
})
}
tresult.Text = result.Text
tresult.Language = result.Language
tresult.Duration = result.Duration
tresult.Eou = result.Eou
return tresult, nil
}
func (s *server) AudioTranscriptionStream(in *pb.TranscriptRequest, stream pb.Backend_AudioTranscriptionStreamServer) error {
if err := s.checkModelIdentity(in); err != nil {
return err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
resultChan := make(chan *pb.TranscriptStreamResponse)
done := make(chan bool)
go func() {
for chunk := range resultChan {
stream.Send(chunk)
}
done <- true
}()
err := s.llm.AudioTranscriptionStream(stream.Context(), in, resultChan)
<-done
return err
}
// AudioTranscriptionLive is the bidirectional live ASR handler. The shape
// mirrors AudioTransformStream exactly (recv → in chan, out chan → send) so
// backends implement it with the same goroutine idiom.
func (s *server) AudioTranscriptionLive(stream pb.Backend_AudioTranscriptionLiveServer) error {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
in := make(chan *pb.TranscriptLiveRequest, 4)
out := make(chan *pb.TranscriptLiveResponse, 4)
// Pump incoming messages from the gRPC stream into `in`. EOF closes the
// channel, which signals the backend to finalize the decode session.
recvErrCh := make(chan error, 1)
go func() {
defer close(in)
for {
req, err := stream.Recv()
if err != nil {
if errors.Is(err, io.EOF) {
recvErrCh <- nil
return
}
recvErrCh <- err
return
}
select {
case in <- req:
case <-stream.Context().Done():
recvErrCh <- stream.Context().Err()
return
}
}
}()
// Pump outgoing responses from `out` to the gRPC stream. The backend
// closes `out` on completion.
sendDone := make(chan error, 1)
go func() {
for resp := range out {
if err := stream.Send(resp); err != nil {
sendDone <- err
// Drain `out` so the backend can finish.
for range out {
}
return
}
}
sendDone <- nil
}()
backendErr := s.llm.AudioTranscriptionLive(in, out)
sendErr := <-sendDone
// Unlike AudioTransformStream, do NOT wait for the recv pump when the
// backend failed: callers block on the first Recv for the ready ack, so
// an unsupported backend (Unimplemented) must surface immediately, not
// after the client gives up and closes its send side. Returning cancels
// the stream context, which unwinds the recv goroutine.
if backendErr != nil {
return backendErr
}
if sendErr != nil {
return sendErr
}
return <-recvErrCh
}
func (s *server) PredictStream(in *pb.PredictOptions, stream pb.Backend_PredictStreamServer) error {
if err := s.checkModelIdentity(in); err != nil {
return err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
if rich, ok := s.llm.(AIModelRich); ok {
replyChan := make(chan *pb.Reply)
done := make(chan bool)
go func() {
for reply := range replyChan {
// Send errors here mean the client disconnected;
// drain the rest of the channel so the producer
// (PredictStreamRich) doesn't block on the next
// reply forever.
_ = stream.Send(reply)
}
done <- true
}()
// Server-side close: PredictStreamRich implementations send into
// the channel and return when finished; closing is the host's
// concern so impls don't have to remember `defer close(...)`.
err := rich.PredictStreamRich(in, replyChan)
close(replyChan)
<-done
return err
}
resultChan := make(chan string)
done := make(chan bool)
go func() {
for result := range resultChan {
stream.Send(newReply(result))
}
done <- true
}()
err := s.llm.PredictStream(in, resultChan)
<-done
return err
}
func (s *server) TokenizeString(ctx context.Context, in *pb.PredictOptions) (*pb.TokenizationResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.TokenizeString(in)
if err != nil {
return nil, err
}
castTokens := make([]int32, len(res.Tokens))
for i, v := range res.Tokens {
castTokens[i] = int32(v)
}
return &pb.TokenizationResponse{
Length: int32(res.Length),
Tokens: castTokens,
}, err
}
func (s *server) Status(ctx context.Context, in *pb.HealthMessage) (*pb.StatusResponse, error) {
res, err := s.llm.Status()
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) StoresSet(ctx context.Context, in *pb.StoresSetOptions) (*pb.Result, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.StoresSet(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error setting entry: %s", err.Error()), Success: false}, err
}
return &pb.Result{Message: "Set key", Success: true}, nil
}
func (s *server) StoresDelete(ctx context.Context, in *pb.StoresDeleteOptions) (*pb.Result, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.StoresDelete(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error deleting entry: %s", err.Error()), Success: false}, err
}
return &pb.Result{Message: "Deleted key", Success: true}, nil
}
func (s *server) StoresGet(ctx context.Context, in *pb.StoresGetOptions) (*pb.StoresGetResult, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.StoresGet(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) StoresFind(ctx context.Context, in *pb.StoresFindOptions) (*pb.StoresFindResult, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.StoresFind(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) VAD(ctx context.Context, in *pb.VADRequest) (*pb.VADResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.VAD(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) Diarize(ctx context.Context, in *pb.DiarizeRequest) (*pb.DiarizeResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.Diarize(in)
if err != nil {
return nil, err
}
return &res, nil
}
func (s *server) SoundDetection(ctx context.Context, in *pb.SoundDetectionRequest) (*pb.SoundDetectionResponse, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
return s.llm.SoundDetection(ctx, in)
}
func (s *server) AudioEncode(ctx context.Context, in *pb.AudioEncodeRequest) (*pb.AudioEncodeResult, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.AudioEncode(in)
if err != nil {
return nil, err
}
return res, nil
}
func (s *server) AudioDecode(ctx context.Context, in *pb.AudioDecodeRequest) (*pb.AudioDecodeResult, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.AudioDecode(in)
if err != nil {
return nil, err
}
return res, nil
}
func (s *server) AudioTransform(ctx context.Context, in *pb.AudioTransformRequest) (*pb.AudioTransformResult, error) {
if err := s.checkModelIdentity(in); err != nil {
return nil, err
}
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.AudioTransform(in)
if err != nil {
return nil, err
}
return res, nil
}
func (s *server) AudioTransformStream(stream pb.Backend_AudioTransformStreamServer) error {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
in := make(chan *pb.AudioTransformFrameRequest, 4)
out := make(chan *pb.AudioTransformFrameResponse, 4)
// Pump incoming frames from the gRPC stream into `in`. EOF closes the
// channel, which signals the backend that the client is done sending.
recvErrCh := make(chan error, 1)
go func() {
defer close(in)
for {
req, err := stream.Recv()
if err != nil {
if errors.Is(err, io.EOF) {
recvErrCh <- nil
return
}
recvErrCh <- err
return
}
select {
case in <- req:
case <-stream.Context().Done():
recvErrCh <- stream.Context().Err()
return
}
}
}()
// Pump outgoing frames from `out` to the gRPC stream. The backend closes
// `out` on completion.
sendDone := make(chan error, 1)
go func() {
for resp := range out {
if err := stream.Send(resp); err != nil {
sendDone <- err
// Drain `out` so the backend can finish.
for range out {
}
return
}
}
sendDone <- nil
}()
backendErr := s.llm.AudioTransformStream(in, out)
sendErr := <-sendDone
recvErr := <-recvErrCh
if backendErr != nil {
return backendErr
}
if sendErr != nil {
return sendErr
}
return recvErr
}
// AudioToAudioStream is the bidirectional any-to-any S2S handler. The
// shape mirrors AudioTransformStream exactly (recv → in chan, out chan →
// send) so backends can implement either via the same goroutine idiom.
func (s *server) AudioToAudioStream(stream pb.Backend_AudioToAudioStreamServer) error {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
in := make(chan *pb.AudioToAudioRequest, 8)
out := make(chan *pb.AudioToAudioResponse, 8)
recvErrCh := make(chan error, 1)
go func() {
defer close(in)
for {
req, err := stream.Recv()
if err != nil {
if errors.Is(err, io.EOF) {
recvErrCh <- nil
return
}
recvErrCh <- err
return
}
select {
case in <- req:
case <-stream.Context().Done():
recvErrCh <- stream.Context().Err()
return
}
}
}()
sendDone := make(chan error, 1)
go func() {
for resp := range out {
if err := stream.Send(resp); err != nil {
sendDone <- err
for range out {
}
return
}
}
sendDone <- nil
}()
backendErr := s.llm.AudioToAudioStream(in, out)
sendErr := <-sendDone
recvErr := <-recvErrCh
if backendErr != nil {
return backendErr
}
if sendErr != nil {
return sendErr
}
return recvErr
}
// Forward is the bidi-stream handler for the cloud-proxy backend's
// passthrough mode. Same recv→in / out→send goroutine idiom as
// AudioTransformStream / AudioToAudioStream above. Buffer size 8 to
// keep SSE token streams flowing — at 4, a half-RTT slow gRPC client
// makes the body-read goroutine in the backend block on out<- after
// every few token frames.
func (s *server) Forward(stream pb.Backend_ForwardServer) error {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
in := make(chan *pb.ForwardRequest, 8)
out := make(chan *pb.ForwardReply, 8)
recvErrCh := make(chan error, 1)
go func() {
defer close(in)
for {
req, err := stream.Recv()
if err != nil {
if errors.Is(err, io.EOF) {
recvErrCh <- nil
return
}
recvErrCh <- err
return
}
select {
case in <- req:
case <-stream.Context().Done():
recvErrCh <- stream.Context().Err()
return
}
}
}()
sendDone := make(chan error, 1)
go func() {
for resp := range out {
if err := stream.Send(resp); err != nil {
sendDone <- err
for range out {
}
return
}
}
sendDone <- nil
}()
backendErr := s.llm.Forward(stream.Context(), in, out)
sendErr := <-sendDone
recvErr := <-recvErrCh
if backendErr != nil {
return backendErr
}
if sendErr != nil {
return sendErr
}
return recvErr
}
func (s *server) StartFineTune(ctx context.Context, in *pb.FineTuneRequest) (*pb.FineTuneJobResult, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.StartFineTune(in)
if err != nil {
return &pb.FineTuneJobResult{Success: false, Message: fmt.Sprintf("Error starting fine-tune: %s", err.Error())}, err
}
return res, nil
}
func (s *server) FineTuneProgress(in *pb.FineTuneProgressRequest, stream pb.Backend_FineTuneProgressServer) error {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
updateChan := make(chan *pb.FineTuneProgressUpdate)
done := make(chan bool)
go func() {
for update := range updateChan {
stream.Send(update)
}
done <- true
}()
err := s.llm.FineTuneProgress(in, updateChan)
<-done
return err
}
func (s *server) StopFineTune(ctx context.Context, in *pb.FineTuneStopRequest) (*pb.Result, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.StopFineTune(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error stopping fine-tune: %s", err.Error()), Success: false}, err
}
return &pb.Result{Message: "Fine-tune stopped", Success: true}, nil
}
func (s *server) ListCheckpoints(ctx context.Context, in *pb.ListCheckpointsRequest) (*pb.ListCheckpointsResponse, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.ListCheckpoints(in)
if err != nil {
return nil, err
}
return res, nil
}
func (s *server) ExportModel(ctx context.Context, in *pb.ExportModelRequest) (*pb.Result, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.ExportModel(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error exporting model: %s", err.Error()), Success: false}, err
}
return &pb.Result{Message: "Model exported", Success: true}, nil
}
func (s *server) StartQuantization(ctx context.Context, in *pb.QuantizationRequest) (*pb.QuantizationJobResult, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.StartQuantization(in)
if err != nil {
return &pb.QuantizationJobResult{Success: false, Message: fmt.Sprintf("Error starting quantization: %s", err.Error())}, err
}
return res, nil
}
func (s *server) QuantizationProgress(in *pb.QuantizationProgressRequest, stream pb.Backend_QuantizationProgressServer) error {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
updateChan := make(chan *pb.QuantizationProgressUpdate)
done := make(chan bool)
go func() {
for update := range updateChan {
stream.Send(update)
}
done <- true
}()
err := s.llm.QuantizationProgress(in, updateChan)
<-done
return err
}
func (s *server) StopQuantization(ctx context.Context, in *pb.QuantizationStopRequest) (*pb.Result, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
err := s.llm.StopQuantization(in)
if err != nil {
return &pb.Result{Message: fmt.Sprintf("Error stopping quantization: %s", err.Error()), Success: false}, err
}
return &pb.Result{Message: "Quantization stopped", Success: true}, nil
}
func (s *server) ModelMetadata(ctx context.Context, in *pb.ModelOptions) (*pb.ModelMetadataResponse, error) {
if s.llm.Locking() {
s.llm.Lock()
defer s.llm.Unlock()
}
res, err := s.llm.ModelMetadata(in)
if err != nil {
return nil, err
}
return res, nil
}
func (s *server) Free(ctx context.Context, in *pb.HealthMessage) (*pb.Result, error) {
if err := s.llm.Free(); err != nil {
return &pb.Result{Success: false, Message: err.Error()}, nil
}
return &pb.Result{Success: true}, nil
}
// NewBackendServer creates a pb.BackendServer.
func NewBackendServer(model AIModel) pb.BackendServer {
return &server{llm: model}
}
// AuthTokenEnvVar is the environment variable used to configure gRPC bearer token auth.
const AuthTokenEnvVar = "LOCALAI_GRPC_AUTH_TOKEN"
// validateToken extracts the bearer token from gRPC metadata and validates it.
func validateToken(ctx context.Context, expected string) error {
md, ok := metadata.FromIncomingContext(ctx)
if !ok {
return status.Error(codes.Unauthenticated, "missing metadata")
}
values := md.Get("authorization")
if len(values) == 0 {
return status.Error(codes.Unauthenticated, "missing authorization header")
}
raw := values[0]
if !strings.HasPrefix(raw, "Bearer ") {
return status.Error(codes.Unauthenticated, "authorization must use Bearer scheme")
}
token := strings.TrimPrefix(raw, "Bearer ")
if subtle.ConstantTimeCompare([]byte(token), []byte(expected)) != 1 {
return status.Error(codes.Unauthenticated, "invalid token")
}
return nil
}
func tokenUnaryInterceptor(token string) grpc.UnaryServerInterceptor {
return func(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
if err := validateToken(ctx, token); err != nil {
return nil, err
}
return handler(ctx, req)
}
}
func tokenStreamInterceptor(token string) grpc.StreamServerInterceptor {
return func(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
if err := validateToken(ss.Context(), token); err != nil {
return err
}
return handler(srv, ss)
}
}
// serverOpts returns the common gRPC server options, including auth interceptors
// when LOCALAI_GRPC_AUTH_TOKEN is set.
func serverOpts() []grpc.ServerOption {
opts := []grpc.ServerOption{
grpc.MaxRecvMsgSize(maxGRPCMessageSize),
grpc.MaxSendMsgSize(maxGRPCMessageSize),
}
if token := os.Getenv(AuthTokenEnvVar); token != "" {
opts = append(opts,
grpc.UnaryInterceptor(tokenUnaryInterceptor(token)),
grpc.StreamInterceptor(tokenStreamInterceptor(token)),
)
log.Printf("gRPC auth enabled via %s", AuthTokenEnvVar)
}
return opts
}
func StartServer(address string, model AIModel) error {
lis, err := net.Listen("tcp", address)
if err != nil {
return err
}
s := grpc.NewServer(serverOpts()...)
pb.RegisterBackendServer(s, &server{llm: model})
log.Printf("gRPC Server listening at %v", lis.Addr())
// Safety net: self-terminate if the LocalAI process that spawned this
// backend dies without running its graceful teardown (see parentwatch.go).
startParentDeathWatcher()
if err := s.Serve(lis); err != nil {
return err
}
return nil
}
func RunServer(address string, model AIModel) (func() error, error) {
lis, err := net.Listen("tcp", address)
if err != nil {
return nil, err
}
s := grpc.NewServer(serverOpts()...)
pb.RegisterBackendServer(s, &server{llm: model})
log.Printf("gRPC Server listening at %v", lis.Addr())
// Safety net: self-terminate if the LocalAI process that spawned this
// backend dies without running its graceful teardown (see parentwatch.go).
startParentDeathWatcher()
if err = s.Serve(lis); err != nil {
return func() error {
return lis.Close()
}, err
}
return func() error {
s.GracefulStop()
return nil
}, nil
}