package daemon

import (
	"context"
	"encoding/json"
	"errors"
	"fmt"
	"io"
	"net/http"
	"net/http/httptest"
	"os"
	"strings"
	"testing"

	containerdclient "github.com/openeuler/Conch/internal/adapters/containerd/client"
	containerdhost "github.com/openeuler/Conch/internal/adapters/containerd/host"
	agentprotocol "github.com/openeuler/Conch/internal/agent/protocol"
	"github.com/openeuler/Conch/internal/conchruntime"
	"github.com/openeuler/Conch/internal/runtimeapi"
	"github.com/openeuler/Conch/internal/sandbox"
)

type fakeSnapshotService struct {
	listReq     runtimeapi.ListSnapshotsOptions
	removeReq   runtimeapi.RemoveSnapshotOptions
	infoReq     runtimeapi.SnapshotInfoOptions
	listErr     error
	removeErr   error
	infoErr     error
	removeCalls int
	snapshots   []runtimeapi.SnapshotRecord
	infoResp    runtimeapi.SnapshotRecord
}

type fakeSandboxOps struct {
	createReq      sandbox.CreateRequest
	checkpointReq  sandbox.CheckpointRequest
	suspendReq     sandbox.LifecycleRequest
	resumeReq      sandbox.LifecycleRequest
	deleteReqs     []sandbox.DeleteRequest
	createErr      error
	createCalls    int
	checkpointErr  error
	checkpointResp sandbox.CheckpointResult
	updateReq      sandbox.NetworkUpdateRequest
}

func newSnapshotHandlerServer(svc conchruntime.SnapshotOps) *Daemon {
	runtimeService := conchruntime.New(nil, nil, nil)
	runtimeService.Snapshot = svc
	s := &Daemon{
		router:         http.NewServeMux(),
		runtimeService: runtimeService,
	}
	s.routes()
	return s
}

func (f *fakeSnapshotService) List(_ context.Context, req runtimeapi.ListSnapshotsOptions) ([]runtimeapi.SnapshotRecord, error) {
	f.listReq = req
	if f.listErr != nil {
		return nil, f.listErr
	}
	return f.snapshots, nil
}

func (f *fakeSnapshotService) Remove(_ context.Context, req runtimeapi.RemoveSnapshotOptions) error {
	f.removeCalls++
	f.removeReq = req
	return f.removeErr
}

func (f *fakeSnapshotService) Info(_ context.Context, req runtimeapi.SnapshotInfoOptions) (runtimeapi.SnapshotRecord, error) {
	f.infoReq = req
	if f.infoErr != nil {
		return runtimeapi.SnapshotRecord{}, f.infoErr
	}
	if f.infoResp.Key != "" {
		return f.infoResp, nil
	}
	return runtimeapi.SnapshotRecord{
		Key:         req.Key,
		Parent:      "parent-id",
		StoragePath: "/snap/rootfs",
	}, nil
}

func (f *fakeSandboxOps) Create(_ context.Context, req sandbox.CreateRequest) (sandbox.CreateResult, error) {
	f.createCalls++
	f.createReq = req
	if f.createErr != nil {
		return sandbox.CreateResult{}, f.createErr
	}
	return sandbox.CreateResult{
		IP:              "192.0.2.2",
		AgentToken:      req.AgentToken,
		SandboxID:       req.SandboxID,
		BootIndexDigest: req.TemplateID,
	}, nil
}

func (f *fakeSandboxOps) Delete(req sandbox.DeleteRequest) error {
	f.deleteReqs = append(f.deleteReqs, req)
	return nil
}

func (f *fakeSandboxOps) Suspend(req sandbox.LifecycleRequest) error {
	f.suspendReq = req
	return nil
}

func (f *fakeSandboxOps) Resume(req sandbox.LifecycleRequest) error {
	f.resumeReq = req
	return nil
}

func (f *fakeSandboxOps) UpdateNetwork(_ context.Context, req sandbox.NetworkUpdateRequest) error {
	f.updateReq = req
	return nil
}

func (f *fakeSandboxOps) Checkpoint(req sandbox.CheckpointRequest) (sandbox.CheckpointResult, error) {
	f.checkpointReq = req
	if f.checkpointErr != nil {
		return sandbox.CheckpointResult{}, f.checkpointErr
	}
	if f.checkpointResp.MemRootPath != "" {
		return f.checkpointResp, nil
	}
	memRoot, err := os.MkdirTemp("", "conch-daemon-checkpoint-test-*")
	if err != nil {
		return sandbox.CheckpointResult{}, err
	}
	return sandbox.CheckpointResult{
		MemRootPath: memRoot,
		VMMName:     "cloud-hypervisor",
	}, nil
}

func TestHandleHealth(t *testing.T) {
	store := newMemorySandboxStore()
	ready := &Daemon{
		sandboxStore:   store,
		containerdHost: &containerdhost.Host{},
		daemonClient:   &containerdclient.Client{},
		runtimeService: &conchruntime.Service{Sandbox: &fakeSandboxOps{}, Store: store},
	}
	for _, test := range []struct {
		name     string
		daemon   *Daemon
		method   string
		want     int
		wantCode string
	}{
		{name: "not ready", daemon: &Daemon{}, method: http.MethodGet, want: http.StatusServiceUnavailable, wantCode: "service.unavailable"},
		{name: "ready", daemon: ready, method: http.MethodGet, want: http.StatusNoContent},
		{name: "method not allowed", daemon: ready, method: http.MethodPost, want: http.StatusMethodNotAllowed, wantCode: "request.method_not_allowed"},
	} {
		t.Run(test.name, func(t *testing.T) {
			recorder := httptest.NewRecorder()
			test.daemon.handleHealth(recorder, httptest.NewRequest(test.method, "/health", nil))
			if recorder.Code != test.want {
				t.Fatalf("status = %d, want %d", recorder.Code, test.want)
			}
			if test.wantCode == "" {
				if recorder.Body.Len() != 0 {
					t.Fatalf("body = %q, want empty", recorder.Body.String())
				}
				return
			}
			var apiErr apiErrorResponse
			if err := json.NewDecoder(recorder.Body).Decode(&apiErr); err != nil {
				t.Fatalf("decode response: %v", err)
			}
			if string(apiErr.Code) != test.wantCode {
				t.Fatalf("code = %q, want %q", apiErr.Code, test.wantCode)
			}
		})
	}
}

func TestMatchesSandboxState(t *testing.T) {
	for _, test := range []struct {
		state  sandbox.State
		states map[string]bool
		want   bool
	}{
		{state: sandbox.StateReady, want: true},
		{state: sandbox.StateSuspended, want: true},
		{state: sandbox.StateUnknown, want: false},
		{state: sandbox.StateSuspended, states: map[string]bool{"paused": true}, want: true},
		{state: sandbox.StateSuspended, states: map[string]bool{"running": true}, want: false},
		{state: sandbox.StateReady, states: map[string]bool{"running": true}, want: true},
		{state: sandbox.StateReady, states: map[string]bool{"paused": true}, want: false},
	} {
		if got := matchesSandboxState(sandbox.Record{State: test.state}, test.states); got != test.want {
			t.Fatalf("matchesSandboxState(%q, %v) = %v, want %v", test.state, test.states, got, test.want)
		}
	}
}

func TestParseSandboxStates(t *testing.T) {
	got, err := parseSandboxStates([]string{"running", "paused"})
	if err != nil || !got["running"] || !got["paused"] {
		t.Fatalf("parseSandboxStates() = %v, %v", got, err)
	}
	if _, err := parseSandboxStates([]string{"stopped"}); err == nil {
		t.Fatal("parseSandboxStates() accepted unsupported state")
	}
	if got, err := parseSandboxStates(nil); err != nil || got != nil {
		t.Fatalf("parseSandboxStates(nil) = %v, %v", got, err)
	}
}

func TestParseSandboxListLimit(t *testing.T) {
	for _, test := range []struct {
		raw  string
		want int
		ok   bool
	}{
		{raw: "", want: 100, ok: true},
		{raw: "1", want: 1, ok: true},
		{raw: "5000", want: 5000, ok: true},
		{raw: "0"},
		{raw: "5001"},
		{raw: "invalid"},
	} {
		got, err := parseSandboxListLimit(test.raw)
		if test.ok && (err != nil || got != test.want) {
			t.Fatalf("parseSandboxListLimit(%q) = %d, %v; want %d", test.raw, got, err, test.want)
		}
		if !test.ok && err == nil {
			t.Fatalf("parseSandboxListLimit(%q) unexpectedly succeeded", test.raw)
		}
	}
}

func TestSandboxV1Handlers(t *testing.T) {
	store := newMemorySandboxStore()

	sandboxOps := &fakeSandboxOps{}
	runtimeService := conchruntime.New(sandboxOps, nil, store)
	runtimeService.SetSandboxDefaults(runtimeapi.SandboxDefaults{
		TemplateName: testTemplateNameDefault,
		VCPUNum:      4,
		VCPUMax:      4,
		RamMB:        256,
	})
	runtimeService.Templates = testTemplateStore()
	server := &Daemon{
		router:         http.NewServeMux(),
		sandboxStore:   store,
		runtimeService: runtimeService,
	}
	server.routes()

	if _, err := store.Create(context.Background(), sandbox.Record{
		ID:                       "sandbox-1",
		State:                    sandbox.StateReady,
		SourceTemplateName:       testTemplateNameExplicit,
		SourceTemplateID:         testTemplateIDExplicit,
		CheckpointHeadTemplateID: "sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa",
		VCPUNum:                  2,
		RamMB:                    128,
	}); err != nil {
		t.Fatalf("seed sandbox: %v", err)
	}

	t.Run("list", func(t *testing.T) {
		response := serveSandboxRequest(server, http.MethodGet, "/api/v1/sandboxes?limit=1", nil)
		if response.Code != http.StatusOK {
			t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
		}
		var records []sandboxInspectResponse
		if err := json.NewDecoder(response.Body).Decode(&records); err != nil {
			t.Fatalf("decode list response: %v", err)
		}
		if len(records) != 1 || records[0].SandboxID != "sandbox-1" || records[0].TemplateName != testTemplateNameExplicit || records[0].TemplateID != testTemplateIDExplicit {
			t.Fatalf("list response = %#v", records)
		}
	})

	t.Run("rejects invalid list queries", func(t *testing.T) {
		for _, path := range []string{
			"/api/v1/sandboxes?state=stopped",
			"/api/v1/sandboxes?limit=0",
			"/api/v1/sandboxes?limit=5001",
		} {
			response := serveSandboxRequest(server, http.MethodGet, path, nil)
			if response.Code != http.StatusBadRequest {
				t.Fatalf("%s status = %d, body = %s", path, response.Code, response.Body.String())
			}
		}
	})

	t.Run("get", func(t *testing.T) {
		response := serveSandboxRequest(server, http.MethodGet, "/api/v1/sandboxes/sandbox-1", nil)
		if response.Code != http.StatusOK {
			t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
		}
		var record sandboxInspectResponse
		if err := json.NewDecoder(response.Body).Decode(&record); err != nil {
			t.Fatalf("decode get response: %v", err)
		}
		if record.SandboxID != "sandbox-1" || record.TemplateName != testTemplateNameExplicit || record.TemplateID != testTemplateIDExplicit || record.Domain == nil {
			t.Fatalf("get response = %#v", record)
		}
	})

	t.Run("create", func(t *testing.T) {
		response := serveSandboxRequest(server, http.MethodPost, "/api/v1/sandboxes", strings.NewReader(`{
			"sandbox_id":"sandbox-2","template_name":"`+testTemplateNameOther+`","env":{"SOME_RANDOM_KEY":"key123"},
			"network":{"denyOut":["192.0.2.10"],"allowIn":["198.51.100.0/24"]}
		}`))
		if response.Code != http.StatusOK {
			t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
		}
		var record createSandboxResponse
		if err := json.NewDecoder(response.Body).Decode(&record); err != nil {
			t.Fatalf("decode create response: %v", err)
		}
		if record.SandboxID != "sandbox-2" || record.TemplateName != testTemplateNameOther || record.TemplateID != testTemplateIDOther ||
			record.Domain != "192.0.2.2" || record.ConchInitAccessToken == "" {
			t.Fatalf("create response = %#v", record)
		}
		if got := sandboxOps.createReq.Env["SOME_RANDOM_KEY"]; got != "key123" {
			t.Fatalf("Env[SOME_RANDOM_KEY] = %q, want key123", got)
		}
		if sandboxOps.createReq.Network == nil || len(sandboxOps.createReq.Network.DenyOut) != 1 || len(sandboxOps.createReq.Network.AllowIn) != 1 {
			t.Fatalf("create network = %#v", sandboxOps.createReq.Network)
		}
	})

	t.Run("rejects invalid environment", func(t *testing.T) {
		for _, test := range []struct {
			body     string
			wantCode string
		}{
			{body: `{"sandbox_id":"invalid-env-key","template_name":"` + testTemplateNameOther + `","env":{"BAD=KEY":"value"}}`, wantCode: "sandbox.invalid_environment"},
			{body: `{"sandbox_id":"invalid-env-value","template_name":"` + testTemplateNameOther + `","env":{"KEY":123}}`, wantCode: "request.invalid_body"},
		} {
			createCalls := sandboxOps.createCalls
			response := serveSandboxRequest(server, http.MethodPost, "/api/v1/sandboxes", strings.NewReader(test.body))
			if response.Code != http.StatusBadRequest {
				t.Fatalf("body = %s, status = %d, response = %s", test.body, response.Code, response.Body.String())
			}
			var apiErr apiErrorResponse
			if err := json.Unmarshal(response.Body.Bytes(), &apiErr); err != nil || string(apiErr.Code) != test.wantCode {
				t.Fatalf("body = %s, decoded response = %#v, error = %v", test.body, apiErr, err)
			}
			if sandboxOps.createCalls != createCalls {
				t.Fatalf("body = %s, runtime Create() calls = %d, want %d", test.body, sandboxOps.createCalls, createCalls)
			}
		}
	})

	t.Run("maps oversized initialization payload to bad request", func(t *testing.T) {
		sandboxOps.createErr = fmt.Errorf("marshal initialization: %w", agentprotocol.ErrPayloadTooLarge)
		t.Cleanup(func() { sandboxOps.createErr = nil })
		response := serveSandboxRequest(server, http.MethodPost, "/api/v1/sandboxes", strings.NewReader(`{
			"sandbox_id":"oversized-env","template_name":"`+testTemplateNameOther+`"
		}`))
		if response.Code != http.StatusBadRequest {
			t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
		}
		var apiErr apiErrorResponse
		if err := json.Unmarshal(response.Body.Bytes(), &apiErr); err != nil || apiErr.Code != "sandbox.initialization_too_large" {
			t.Fatalf("decoded response = %#v, error = %v", apiErr, err)
		}
		sandboxOps.createErr = nil
	})

	t.Run("update network", func(t *testing.T) {
		response := serveSandboxRequest(server, http.MethodPut, "/api/v1/sandboxes/sandbox-1/network", strings.NewReader(`{
			"allowOut":["192.0.2.20"],"denyIn":["198.51.100.20"]
		}`))
		if response.Code != http.StatusNoContent {
			t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
		}
		if sandboxOps.updateReq.Network == nil || len(sandboxOps.updateReq.Network.AllowOut) != 1 || len(sandboxOps.updateReq.Network.DenyIn) != 1 {
			t.Fatalf("update network = %#v", sandboxOps.updateReq.Network)
		}
		getResponse := serveSandboxRequest(server, http.MethodGet, "/api/v1/sandboxes/sandbox-1", nil)
		var record sandboxInspectResponse
		if getResponse.Code != http.StatusOK || json.NewDecoder(getResponse.Body).Decode(&record) != nil {
			t.Fatalf("get after update status = %d, body = %s", getResponse.Code, getResponse.Body.String())
		}
		if record.Network == nil || len(record.Network.AllowOut) != 1 || len(record.Network.DenyIn) != 1 {
			t.Fatalf("get network = %#v", record.Network)
		}
	})

	t.Run("rejects invalid network", func(t *testing.T) {
		response := serveSandboxRequest(server, http.MethodPut, "/api/v1/sandboxes/sandbox-1/network", strings.NewReader(`{"allowOut":["example.com"]}`))
		if response.Code != http.StatusBadRequest {
			t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
		}
		response = serveSandboxRequest(server, http.MethodPost, "/api/v1/sandboxes", strings.NewReader(`{
			"sandbox_id":"invalid-network","template_name":"`+testTemplateNameOther+`","network":{"denyIn":["example.com"]}
		}`))
		if response.Code != http.StatusBadRequest {
			t.Fatalf("create status = %d, body = %s", response.Code, response.Body.String())
		}
	})

	t.Run("rejects unknown sandbox subroute", func(t *testing.T) {
		response := serveSandboxRequest(server, http.MethodGet, "/api/v1/sandboxes/sandbox-1/unknown", nil)
		if response.Code != http.StatusNotFound {
			t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
		}
	})

	t.Run("delete", func(t *testing.T) {
		response := serveSandboxRequest(server, http.MethodDelete, "/api/v1/sandboxes/sandbox-1", nil)
		if response.Code != http.StatusNoContent {
			t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
		}
		if len(sandboxOps.deleteReqs) != 1 || sandboxOps.deleteReqs[0].SandboxID != "sandbox-1" {
			t.Fatalf("delete requests = %#v", sandboxOps.deleteReqs)
		}
		if _, err := store.Get(context.Background(), "sandbox-1"); !errors.Is(err, sandbox.ErrNotFound) {
			t.Fatalf("deleted sandbox lookup error = %v", err)
		}
	})
}

func serveSandboxRequest(server *Daemon, method, path string, body io.Reader) *httptest.ResponseRecorder {
	recorder := httptest.NewRecorder()
	server.router.ServeHTTP(recorder, httptest.NewRequest(method, path, body))
	return recorder
}