Restrict HTTP file artifact inputs
This commit is contained in:
@@ -7,6 +7,7 @@ import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"gitea.maximumdirect.net/eric/scriptorium/internal/artifact"
|
||||
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/scriptorium/internal/profile"
|
||||
"gitea.maximumdirect.net/eric/scriptorium/internal/promptdef"
|
||||
@@ -187,6 +188,8 @@ func mapRunError(err error) (int, string, string) {
|
||||
return http.StatusBadRequest, "prompt_load_failed", "failed to load prompt definition"
|
||||
case errors.Is(err, usecase.ErrProfileLoad):
|
||||
return http.StatusBadRequest, "profile_load_failed", "failed to load execution profile"
|
||||
case errors.Is(err, artifact.ErrFileNotAllowed), errors.Is(err, artifact.ErrFileOutsideRoot):
|
||||
return http.StatusBadRequest, "artifact_not_allowed", "file input artifact is not allowed"
|
||||
case errors.Is(err, usecase.ErrArtifactLoad):
|
||||
return http.StatusBadRequest, "artifact_read_failed", "failed to read input artifact"
|
||||
case errors.Is(err, usecase.ErrPromptRender):
|
||||
|
||||
@@ -7,11 +7,14 @@ import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gitea.maximumdirect.net/eric/scriptorium/internal/artifact"
|
||||
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
|
||||
"gitea.maximumdirect.net/eric/scriptorium/internal/llm"
|
||||
"gitea.maximumdirect.net/eric/scriptorium/internal/profile"
|
||||
@@ -61,6 +64,12 @@ func (handlerRenderer) Render(ctx context.Context, definition *domain.PromptDefi
|
||||
return &domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, nil
|
||||
}
|
||||
|
||||
type handlerLLMClient struct{}
|
||||
|
||||
func (handlerLLMClient) Generate(ctx context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) {
|
||||
return &domain.GenerateResponse{Content: "ok"}, nil
|
||||
}
|
||||
|
||||
func TestHandlerPostRunsSuccessWithExplicitProfileID(t *testing.T) {
|
||||
start := time.Now().UTC()
|
||||
end := start.Add(2 * time.Second)
|
||||
@@ -184,6 +193,87 @@ func TestHandlerPostRunsSuccessWithExplicitProfileID(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerInlineRefsWorkWithoutArtifactRoot(t *testing.T) {
|
||||
h := newArtifactRootHandler(t, "")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/runs", bytes.NewBufferString(`{
|
||||
"prompt_id":"p",
|
||||
"inputs":{"x":{"type":"inline","body":"inline body"}}
|
||||
}`))
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
h.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200, got %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerFileRefsWithoutArtifactRootAreRejected(t *testing.T) {
|
||||
h := newArtifactRootHandler(t, "")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/runs", bytes.NewBufferString(`{
|
||||
"prompt_id":"p",
|
||||
"inputs":{"x":{"type":"file","uri":"input.txt"}}
|
||||
}`))
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
h.ServeHTTP(w, req)
|
||||
|
||||
assertHTTPErrorCode(t, w, http.StatusBadRequest, "artifact_not_allowed")
|
||||
}
|
||||
|
||||
func TestHandlerFileRefsUnderArtifactRootWork(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(root, "input.txt"), []byte("allowed"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
h := newArtifactRootHandler(t, root)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/runs", bytes.NewBufferString(`{
|
||||
"prompt_id":"p",
|
||||
"inputs":{"x":{"type":"file","uri":"input.txt"}}
|
||||
}`))
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
h.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200, got %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerFileRefsOutsideArtifactRootAreRejected(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
outside := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(outside, "secret.txt"), []byte("denied"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
h := newArtifactRootHandler(t, root)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
uri string
|
||||
}{
|
||||
{name: "relative traversal", uri: filepath.Join("..", filepath.Base(outside), "secret.txt")},
|
||||
{name: "absolute outside root", uri: filepath.Join(outside, "secret.txt")},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
body := fmt.Sprintf(`{
|
||||
"prompt_id":"p",
|
||||
"inputs":{"x":{"type":"file","uri":%q}}
|
||||
}`, tc.uri)
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/runs", bytes.NewBufferString(body))
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
h.ServeHTTP(w, req)
|
||||
|
||||
assertHTTPErrorCode(t, w, http.StatusBadRequest, "artifact_not_allowed")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerPostRunsSuccessUsingPromptDefaultProfile(t *testing.T) {
|
||||
r := &fakeRunner{result: &domain.RunResult{
|
||||
Artifact: domain.Artifact{Body: []byte("ok")},
|
||||
@@ -676,3 +766,48 @@ func TestHandlerValidationFailureStillSuccessAndRawOutputOptIn(t *testing.T) {
|
||||
func wrap(stage error, cause error) error {
|
||||
return fmt.Errorf("%w: %w", stage, cause)
|
||||
}
|
||||
|
||||
func newArtifactRootHandler(t *testing.T, root string) *Handler {
|
||||
t.Helper()
|
||||
|
||||
reader, err := artifact.NewRestrictedCompositeReader(root)
|
||||
if err != nil {
|
||||
t.Fatalf("expected restricted artifact reader: %v", err)
|
||||
}
|
||||
runner := usecase.NewRunner(
|
||||
handlerPromptRepo{def: &domain.PromptDefinition{
|
||||
ID: "p",
|
||||
Version: "1",
|
||||
DefaultProfile: "exec",
|
||||
Templates: []domain.PromptMessageTemplate{{Role: "user", Content: "hi"}},
|
||||
OutputFormat: domain.FormatText,
|
||||
Validation: domain.OutputContract{Format: domain.FormatText, ValidationMode: domain.ValidationNone},
|
||||
}},
|
||||
handlerProfileRepo{profile: &domain.ExecutionProfile{
|
||||
ID: "exec",
|
||||
Endpoint: "http://example.invalid/v1",
|
||||
Model: "model",
|
||||
}},
|
||||
reader,
|
||||
handlerRenderer{},
|
||||
handlerLLMClient{},
|
||||
nil,
|
||||
)
|
||||
return NewHandler(runner)
|
||||
}
|
||||
|
||||
func assertHTTPErrorCode(t *testing.T, w *httptest.ResponseRecorder, status int, code string) {
|
||||
t.Helper()
|
||||
|
||||
if w.Code != status {
|
||||
t.Fatalf("expected %d, got %d body=%s", status, w.Code, w.Body.String())
|
||||
}
|
||||
var resp map[string]any
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("invalid JSON response: %v", err)
|
||||
}
|
||||
errBody := resp["error"].(map[string]any)
|
||||
if errBody["code"] != code {
|
||||
t.Fatalf("expected code %q, got %#v", code, errBody["code"])
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user