Files
chorus/internal/platform/mockprovider/handler.go
T

134 lines
3.7 KiB
Go

package mockprovider
import (
"bytes"
"encoding/base64"
"encoding/json"
"image"
"image/color"
"image/png"
"net/http"
"strings"
"time"
)
type Handler struct{ Delay time.Duration }
func (h Handler) ServeHTTP(response http.ResponseWriter, request *http.Request) {
switch request.URL.Path {
case "/v1/chat/completions":
h.chat(response, request)
case "/v1/images/edits":
h.image(response, request)
case "/v1/result.png":
response.Header().Set("Content-Type", "image/png")
response.Write(mockPNG())
default:
http.NotFound(response, request)
}
}
func (h Handler) chat(response http.ResponseWriter, request *http.Request) {
var body struct {
Messages []struct {
Content string `json:"content"`
} `json:"messages"`
}
if json.NewDecoder(http.MaxBytesReader(response, request.Body, 1<<20)).Decode(&body) != nil || len(body.Messages) == 0 {
writeError(response, 400, "bad_request")
return
}
prompt := body.Messages[0].Content
if h.scenario(response, request, prompt) {
return
}
response.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(response).Encode(map[string]any{"choices": []any{map[string]any{"message": map[string]string{"content": "mock: " + prompt}}}})
}
func (h Handler) image(response http.ResponseWriter, request *http.Request) {
if request.ParseMultipartForm(8<<20) != nil {
writeError(response, 400, "bad_request")
return
}
prompt := request.FormValue("prompt")
if h.scenario(response, request, prompt) {
return
}
files := request.MultipartForm.File["image[]"]
if len(files) == 0 {
writeError(response, 400, "bad_request")
return
}
for _, header := range files {
file, err := header.Open()
if err != nil {
writeError(response, 400, "bad_request")
return
}
_, err = png.Decode(file)
file.Close()
if err != nil {
writeError(response, 400, "bad_image")
return
}
}
data := map[string]any{"b64_json": base64.StdEncoding.EncodeToString(mockPNG())}
if strings.Contains(prompt, "mock:url") {
data = map[string]any{"url": "http://" + request.Host + "/v1/result.png"}
}
response.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(response).Encode(map[string]any{"data": []any{data}})
}
func (h Handler) scenario(response http.ResponseWriter, request *http.Request, prompt string) bool {
switch {
case strings.Contains(prompt, "mock:429"):
writeError(response, 429, "rate_limited")
case strings.Contains(prompt, "mock:500"):
writeError(response, 500, "server_error")
case strings.Contains(prompt, "mock:400"):
writeError(response, 400, "bad_request")
case strings.Contains(prompt, "mock:401"):
writeError(response, 401, "unauthorized")
case strings.Contains(prompt, "mock:policy"):
writeError(response, 400, "content_policy_violation")
case strings.Contains(prompt, "mock:timeout"):
delay := h.Delay
if delay <= 0 {
delay = time.Second
}
select {
case <-request.Context().Done():
case <-time.After(delay):
writeError(response, 504, "timeout")
}
case strings.Contains(prompt, "mock:connection"):
if hijacker, ok := response.(http.Hijacker); ok {
connection, _, err := hijacker.Hijack()
if err == nil {
connection.Close()
}
}
default:
return false
}
return true
}
func writeError(response http.ResponseWriter, status int, code string) {
response.Header().Set("Content-Type", "application/json")
response.WriteHeader(status)
_ = json.NewEncoder(response).Encode(map[string]any{"error": map[string]string{"code": code, "message": "mock upstream failure"}})
}
func mockPNG() []byte {
img := image.NewRGBA(image.Rect(0, 0, 2, 2))
for y := 0; y < 2; y++ {
for x := 0; x < 2; x++ {
img.Set(x, y, color.RGBA{R: 40, G: 120, B: 200, A: 255})
}
}
var output bytes.Buffer
_ = png.Encode(&output, img)
return output.Bytes()
}