191 lines
5.6 KiB
Go
191 lines
5.6 KiB
Go
package mockprovider
|
|
|
|
import (
|
|
"bytes"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"image"
|
|
"image/color"
|
|
"image/png"
|
|
"mime/multipart"
|
|
"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/generations":
|
|
h.image(response, request, false)
|
|
case "/v1/images/edits":
|
|
h.image(response, request, true)
|
|
case "/v1/result.png":
|
|
response.Header().Set("Content-Type", "image/png")
|
|
response.Write(mockPNG())
|
|
default:
|
|
if strings.HasPrefix(request.URL.Path, "/v1/models/") && strings.HasSuffix(request.URL.Path, ":generateContent") {
|
|
h.gemini(response, request)
|
|
return
|
|
}
|
|
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, requireInput bool) {
|
|
prompt := ""
|
|
var files []*multipart.FileHeader
|
|
if strings.HasPrefix(request.Header.Get("Content-Type"), "application/json") {
|
|
var body struct {
|
|
Prompt string `json:"prompt"`
|
|
}
|
|
if json.NewDecoder(http.MaxBytesReader(response, request.Body, 1<<20)).Decode(&body) != nil {
|
|
writeError(response, 400, "bad_request")
|
|
return
|
|
}
|
|
prompt = body.Prompt
|
|
} else {
|
|
if request.ParseMultipartForm(8<<20) != nil {
|
|
writeError(response, 400, "bad_request")
|
|
return
|
|
}
|
|
prompt = request.FormValue("prompt")
|
|
files = request.MultipartForm.File["image[]"]
|
|
}
|
|
if h.scenario(response, request, prompt) {
|
|
return
|
|
}
|
|
if requireInput && 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) gemini(response http.ResponseWriter, request *http.Request) {
|
|
var body struct {
|
|
Contents []struct {
|
|
Parts []struct {
|
|
Text string `json:"text"`
|
|
InlineData *struct {
|
|
MIMEType string `json:"mimeType"`
|
|
Data string `json:"data"`
|
|
} `json:"inlineData"`
|
|
} `json:"parts"`
|
|
} `json:"contents"`
|
|
}
|
|
if json.NewDecoder(http.MaxBytesReader(response, request.Body, 1<<20)).Decode(&body) != nil || len(body.Contents) != 1 {
|
|
writeError(response, 400, "bad_request")
|
|
return
|
|
}
|
|
prompt := ""
|
|
hasInlineData := false
|
|
for _, part := range body.Contents[0].Parts {
|
|
if part.Text != "" {
|
|
prompt = part.Text
|
|
}
|
|
if part.InlineData != nil && part.InlineData.MIMEType != "" && part.InlineData.Data != "" {
|
|
hasInlineData = true
|
|
}
|
|
}
|
|
if prompt == "" || h.scenario(response, request, prompt) {
|
|
return
|
|
}
|
|
part := map[string]any{"text": "mock: " + prompt}
|
|
if hasInlineData {
|
|
part = map[string]any{"inlineData": map[string]string{"mimeType": "image/png", "data": base64.StdEncoding.EncodeToString(mockPNG())}}
|
|
}
|
|
response.Header().Set("Content-Type", "application/json")
|
|
_ = json.NewEncoder(response).Encode(map[string]any{"candidates": []any{map[string]any{"content": map[string]any{"parts": []any{part}}}}})
|
|
}
|
|
|
|
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()
|
|
}
|