76 lines
1.9 KiB
Go
76 lines
1.9 KiB
Go
package rag
|
|
|
|
import (
|
|
"context"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
|
|
"github.com/google/uuid"
|
|
)
|
|
|
|
type resolverStub struct {
|
|
release Release
|
|
err error
|
|
}
|
|
|
|
func (stub resolverStub) Resolve(
|
|
context.Context,
|
|
uuid.UUID,
|
|
uuid.UUID,
|
|
string,
|
|
) (Release, error) {
|
|
return stub.release, stub.err
|
|
}
|
|
|
|
func TestRetrieveValidatesQuery(t *testing.T) {
|
|
handler := NewHandler(resolverStub{}, UnavailableEngine{})
|
|
request := httptest.NewRequest(http.MethodPost, "/v1/retrieve", strings.NewReader(`{
|
|
"project_id":"00000000-0000-0000-0000-000000000001",
|
|
"knowledge_base_id":"00000000-0000-0000-0000-000000000002",
|
|
"query":""
|
|
}`))
|
|
response := httptest.NewRecorder()
|
|
|
|
handler.retrieve(response, request)
|
|
|
|
if response.Code != http.StatusBadRequest {
|
|
t.Fatalf("expected 400, got %d", response.Code)
|
|
}
|
|
}
|
|
|
|
func TestRetrieveReportsUnavailableBackend(t *testing.T) {
|
|
handler := NewHandler(
|
|
resolverStub{release: Release{ID: uuid.New(), Alias: "production", Version: 1}},
|
|
UnavailableEngine{},
|
|
)
|
|
request := httptest.NewRequest(http.MethodPost, "/v1/retrieve", strings.NewReader(`{
|
|
"project_id":"00000000-0000-0000-0000-000000000001",
|
|
"knowledge_base_id":"00000000-0000-0000-0000-000000000002",
|
|
"query":"example"
|
|
}`))
|
|
response := httptest.NewRecorder()
|
|
|
|
handler.retrieve(response, request)
|
|
|
|
if response.Code != http.StatusServiceUnavailable {
|
|
t.Fatalf("expected 503, got %d", response.Code)
|
|
}
|
|
}
|
|
|
|
func TestRetrieveRejectsTrailingJSON(t *testing.T) {
|
|
handler := NewHandler(resolverStub{}, UnavailableEngine{})
|
|
request := httptest.NewRequest(http.MethodPost, "/v1/retrieve", strings.NewReader(`{
|
|
"project_id":"00000000-0000-0000-0000-000000000001",
|
|
"knowledge_base_id":"00000000-0000-0000-0000-000000000002",
|
|
"query":"example"
|
|
} {}`))
|
|
response := httptest.NewRecorder()
|
|
|
|
handler.retrieve(response, request)
|
|
|
|
if response.Code != http.StatusBadRequest {
|
|
t.Fatalf("expected 400, got %d", response.Code)
|
|
}
|
|
}
|