diff --git a/backend/go/localai-proxy/text.go b/backend/go/localai-proxy/text.go index 2c7a6a0ff..34bc98411 100644 --- a/backend/go/localai-proxy/text.go +++ b/backend/go/localai-proxy/text.go @@ -32,6 +32,13 @@ type textRequest struct { Stop []string `json:"stop,omitempty"` Tools json.RawMessage `json:"tools,omitempty"` ToolChoice json.RawMessage `json:"tool_choice,omitempty"` + // StreamOptions is set on streamed requests: the upstream sends the + // usage trailer only when include_usage asks for it. + StreamOptions *streamOptions `json:"stream_options,omitempty"` +} + +type streamOptions struct { + IncludeUsage bool `json:"include_usage"` } type chatMessage struct { @@ -102,6 +109,9 @@ func (p *LocalAIProxy) textRequest(opts *pb.PredictOptions, stream bool) (string Tools: rawJSON(opts.GetTools()), ToolChoice: rawJSON(opts.GetToolChoice()), } + if stream { + req.StreamOptions = &streamOptions{IncludeUsage: true} + } if len(opts.GetMessages()) == 0 { req.Prompt = opts.GetPrompt() return "/v1/completions", req @@ -298,6 +308,11 @@ func (p *LocalAIProxy) Embeddings(opts *pb.PredictOptions) ([]float32, error) { } `json:"data"` } body := map[string]any{"model": p.model(""), "input": opts.GetEmbeddings()} + // Core sends tokenized input in EmbeddingTokens and leaves Embeddings + // empty; a list of token lists is how the REST API takes tokens. + if tokens := opts.GetEmbeddingTokens(); len(tokens) > 0 { + body["input"] = [][]int32{tokens} + } if err := p.postJSON(context.Background(), "/v1/embeddings", body, &resp); err != nil { return nil, err } diff --git a/backend/go/localai-proxy/text_test.go b/backend/go/localai-proxy/text_test.go index 7887a7d38..e6e2f787a 100644 --- a/backend/go/localai-proxy/text_test.go +++ b/backend/go/localai-proxy/text_test.go @@ -251,6 +251,25 @@ var _ = Describe("localai-proxy", func() { Expect(up.last().JSON).To(HaveKeyWithValue("stream", true)) }) + It("asks the upstream for the usage trailer and reports its token counts", 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{"content": "hi"}}}}), + sseJSON(map[string]any{"choices": []any{}, "usage": map[string]any{"prompt_tokens": 7, "completion_tokens": 1}}), + "[DONE]", + }}) + + results := make(chan *pb.Reply, 10) + Expect(p.PredictStreamRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, results)).To(Succeed()) + // LocalAI (like OpenAI) only sends the usage trailer on request. + Expect(up.last().JSON).To(HaveKeyWithValue("stream_options", HaveKeyWithValue("include_usage", true))) + Expect(results).To(HaveLen(2)) + <-results + usage := <-results + Expect(usage.GetPromptTokens()).To(Equal(int32(7))) + Expect(usage.GetTokens()).To(Equal(int32(1))) + }) + It("streams /v1/completions text for a bare prompt", func() { p := loadProxy(up, nil) up.script("/v1/completions", scriptedResponse{SSE: []string{ @@ -335,6 +354,17 @@ var _ = Describe("localai-proxy", func() { Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) }) + It("sends token input as a token array, not as an empty string", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/embeddings", map[string]any{ + "data": []any{map[string]any{"embedding": []float32{0.5}}}, + }) + + _, err := p.Embeddings(&pb.PredictOptions{EmbeddingTokens: []int32{1, 2, 3}}) + Expect(err).NotTo(HaveOccurred()) + Expect(up.last().JSON).To(HaveKeyWithValue("input", []any{[]any{1.0, 2.0, 3.0}})) + }) + It("fails when the upstream returns no vector", func() { p := loadProxy(up, nil) up.replyJSON("/v1/embeddings", map[string]any{"data": []any{}})