// Package notes provides the notes tool for agent memory with disk persistence.
package notes

import (
	"crypto/sha256"
	"encoding/json"
	"fmt"
	"log"
	"os"
	"path/filepath"
	"strings"
	"sync"

	"github.com/xalgord/xalgorix/v4/internal/sandbox"
	"github.com/xalgord/xalgorix/v4/internal/scanctx"
	"github.com/xalgord/xalgorix/v4/internal/tools"
)

// ── Per-instance note stores ──
var (
	noteStores   = make(map[string]*noteStore)
	noteStoresMu sync.RWMutex
)

type noteStore struct {
	mu sync.RWMutex
	// contextID is the owning ScanContext.ID for this store. Captured at
	// creation time so writeFile can look up the active ScanContext for
	// Path_Policy resolution (sc.ScanDir-relative roots).
	contextID   string
	store       map[string]string
	persistPath string
}

func getNoteStoreByID(id string) *noteStore {
	noteStoresMu.RLock()
	s, ok := noteStores[id]
	noteStoresMu.RUnlock()
	if ok {
		return s
	}

	noteStoresMu.Lock()
	defer noteStoresMu.Unlock()
	if s, ok := noteStores[id]; ok {
		return s
	}
	s = &noteStore{contextID: id, store: make(map[string]string)}
	noteStores[id] = s
	return s
}

// getNoteStore returns the note store for the default (CLI) scan context.
func getNoteStore() *noteStore {
	return getNoteStoreByID(scanctx.Default().ID)
}

func getNoteStoreForContext(contextID string) *noteStore {
	noteStoresMu.RLock()
	s, ok := noteStores[contextID]
	noteStoresMu.RUnlock()
	if ok {
		return s
	}
	noteStoresMu.Lock()
	defer noteStoresMu.Unlock()
	if s, ok := noteStores[contextID]; ok {
		return s
	}
	s = &noteStore{contextID: contextID, store: make(map[string]string)}
	noteStores[contextID] = s
	return s
}

// SetPersistPath configures the directory where notes.json will be saved.
func SetPersistPath(dir string) {
	s := getNoteStore()
	s.mu.Lock()
	defer s.mu.Unlock()
	if dir != "" {
		s.persistPath = filepath.Join(dir, "notes.json")
	} else {
		s.persistPath = ""
	}
}

// ResetNotes clears all notes for the active scan context.
func ResetNotes() {
	s := getNoteStore()
	s.mu.Lock()
	s.store = make(map[string]string)
	s.mu.Unlock()
}

// LoadFromDisk loads notes from the persist path if it exists.
func LoadFromDisk() int {
	s := getNoteStore()
	s.mu.Lock()
	defer s.mu.Unlock()

	if s.persistPath == "" {
		return 0
	}

	data, err := os.ReadFile(s.persistPath)
	if err != nil {
		return 0
	}

	loaded := make(map[string]string)
	if err := json.Unmarshal(data, &loaded); err != nil {
		log.Printf("[notes] Warning: failed to parse %s: %v", s.persistPath, err)
		return 0
	}

	count := 0
	for k, v := range loaded {
		if _, exists := s.store[k]; !exists {
			s.store[k] = v
			count++
		}
	}

	if count > 0 {
		log.Printf("[notes] Loaded %d notes from disk: %s", count, s.persistPath)
	}
	return count
}

// marshalSnapshot serializes the current store. Must be called with s.mu held.
// Returns the serialized data and the persist path. If persistPath is empty,
// returns nil data (meaning no write is needed).
func (s *noteStore) marshalSnapshot() ([]byte, string) {
	if s.persistPath == "" {
		return nil, ""
	}
	data, err := json.MarshalIndent(s.store, "", "  ")
	if err != nil {
		log.Printf("[notes] Warning: failed to marshal notes: %v", err)
		return nil, ""
	}
	return data, s.persistPath
}

// writeFile persists serialized data to disk. Safe to call without holding s.mu.
//
// Every persistence write flows through sandbox.Default().CheckResolve so a
// misconfigured persistPath cannot escape the Allow_List (Requirement 8.3).
// The lookup uses the owning ScanContext's roots when available, falling
// back to the process Workspace_Root via sandbox.Resolve.
func (s *noteStore) writeFile(data []byte, path string) error {
	if data == nil || path == "" {
		return nil
	}
	canonical, err := sandbox.Default().CheckResolve(scanctx.Get(s.contextID), "notes", path)
	if err != nil {
		log.Printf("[notes] Warning: refusing to save notes to %s: %v", path, err)
		return err
	}
	// Atomic persistence: write to a temp file and rename into place so an
	// interrupted write can never leave a truncated/corrupt snapshot at the
	// canonical path (a half-written snapshot would parse as zero notes on
	// restore and silently erase the context's durable note history).
	tmp := canonical + ".tmp"
	if err := os.WriteFile(tmp, data, 0600); err != nil {
		log.Printf("[notes] Warning: failed to save notes to %s: %v", tmp, err)
		return err
	}
	if err := os.Rename(tmp, canonical); err != nil {
		_ = os.Remove(tmp)
		log.Printf("[notes] Warning: failed to save notes to %s: %v", canonical, err)
		return err
	}
	return nil
}

// Register adds note tools to the registry.
func Register(r *tools.Registry) {
	r.Register(&tools.Tool{
		Name:        "add_note",
		Description: "Add a note to persistent memory. Use this to track: discovered endpoints, parameters, tech stack, CSRF tokens, session cookies, exploit chain state, intermediate findings, and anything needed across multiple iterations. Notes persist for the entire scan AND survive context pruning. Use structured keys like 'csrf_token', 'admin_endpoint', 'sqli_confirmed', 'angular_version'.",
		Parameters: []tools.Parameter{
			{Name: "key", Description: "Unique key for the note (e.g., 'csrf_token', 'endpoint_inventory'). Provide it when possible; value-only notes receive a stable key automatically.", Required: false},
			{Name: "value", Description: "Note content", Required: true},
		},
		Execute: func(args map[string]string) (tools.Result, error) {
			return addNoteForContext(r.GetScanContextID(), args)
		},
	})

	r.Register(&tools.Tool{
		Name:        "read_notes",
		Description: "Read all notes or a specific note from memory.",
		Parameters: []tools.Parameter{
			{Name: "key", Description: "Key to read (omit for all notes)", Required: false},
		},
		Execute: func(args map[string]string) (tools.Result, error) {
			return readNotesForContext(r.GetScanContextID(), args)
		},
	})
}

//lint:ignore U1000 kept as a package-level compatibility wrapper for callers in this package.
func addNote(args map[string]string) (tools.Result, error) {
	return addNoteForContext(scanctx.Default().ID, args)
}

func addNoteForContext(contextID string, args map[string]string) (tools.Result, error) {
	key := strings.TrimSpace(args["key"])
	value := strings.TrimSpace(args["value"])
	if value == "" {
		return tools.Result{}, fmt.Errorf("add_note requires a nonempty value")
	}
	if key == "" {
		if lower := strings.ToLower(value); strings.Contains(lower, "endpoint") || strings.Contains(lower, "route inventory") {
			key = "endpoint_inventory"
		} else {
			sum := sha256.Sum256([]byte(value))
			key = fmt.Sprintf("note_%x", sum[:8])
		}
	}

	if contextID == "" {
		contextID = scanctx.Default().ID
	}
	s := getNoteStoreForContext(contextID)
	// The durable write happens UNDER the state lock: snapshots are written in
	// the order they were taken, so a slower older snapshot can no longer
	// overwrite a newer one under concurrency. Map readers are unaffected
	// (they lock the map, not the file).
	s.mu.Lock()
	prev, existed := s.store[key]
	s.store[key] = value
	data, path := s.marshalSnapshot()
	writeErr := s.writeFile(data, path)
	if writeErr != nil {
		// Durable-write acknowledgement: the receipt must never claim a save
		// that failed on disk. Roll the in-memory change back so receipt and
		// durable state agree, and tell the caller what to do about it.
		if existed {
			s.store[key] = prev
		} else {
			delete(s.store, key)
		}
	}
	s.mu.Unlock()

	if writeErr != nil {
		return tools.Result{Error: fmt.Sprintf("add_note: durable write failed (%v) — the note was NOT saved. Retry the add_note call, or include the information in your report instead of relying on it persisting.", writeErr)}, nil
	}
	return tools.Result{Output: fmt.Sprintf("Note saved: %s", key)}, nil
}

//lint:ignore U1000 kept as a package-level compatibility wrapper for callers in this package.
func readNotes(args map[string]string) (tools.Result, error) {
	return readNotesForContext(scanctx.Default().ID, args)
}

func readNotesForContext(contextID string, args map[string]string) (tools.Result, error) {
	key := args["key"]

	if contextID == "" {
		contextID = scanctx.Default().ID
	}
	s := getNoteStoreForContext(contextID)
	s.mu.RLock()
	defer s.mu.RUnlock()

	if key != "" {
		v, ok := s.store[key]
		if !ok {
			return tools.Result{Output: fmt.Sprintf("No note found with key: %s", key)}, nil
		}
		return tools.Result{Output: v}, nil
	}

	if len(s.store) == 0 {
		return tools.Result{Output: "(no notes yet)"}, nil
	}

	var b strings.Builder
	for k, v := range s.store {
		b.WriteString(fmt.Sprintf("📝 %s:\n%s\n\n", k, v))
	}
	return tools.Result{Output: b.String()}, nil
}

// GetAllNotes returns all notes as a map for the active scan context.
func GetAllNotes() map[string]string {
	s := getNoteStore()
	s.mu.RLock()
	defer s.mu.RUnlock()
	result := make(map[string]string, len(s.store))
	for k, v := range s.store {
		result[k] = v
	}
	return result
}

// FormatForContext returns a compact summary of all notes for the active scan context.
func FormatForContext() string {
	s := getNoteStore()
	return formatStore(s)
}

// FormatForContextID returns a compact summary of notes for a specific scan context ID.
func FormatForContextID(contextID string) string {
	s := getNoteStoreForContext(contextID)
	return formatStore(s)
}

func formatStore(s *noteStore) string {
	s.mu.RLock()
	defer s.mu.RUnlock()

	if len(s.store) == 0 {
		return ""
	}

	var b strings.Builder
	b.WriteString("=== YOUR SAVED NOTES (from add_note) ===\n")
	for k, v := range s.store {
		if len(v) > 500 {
			v = v[:500] + "... (truncated)"
		}
		b.WriteString(fmt.Sprintf("• %s: %s\n", k, v))
	}
	b.WriteString("=== END NOTES ===")
	return b.String()
}

// SetPersistPathForContext configures disk persistence for a specific context.
func SetPersistPathForContext(contextID, dir string) {
	s := getNoteStoreForContext(contextID)
	s.mu.Lock()
	defer s.mu.Unlock()
	if dir != "" {
		s.persistPath = filepath.Join(dir, "notes.json")
	} else {
		s.persistPath = ""
	}
}

// ResetNotesForContext clears all notes for a specific context.
func ResetNotesForContext(contextID string) {
	s := getNoteStoreForContext(contextID)
	s.mu.Lock()
	s.store = make(map[string]string)
	s.mu.Unlock()
}

// LoadFromDiskForContext loads notes from disk for a specific context.
// LoadFromDiskForContext loads notes from disk for a specific context. The
// returned error distinguishes a restore failure (missing file, corrupt JSON)
// from a genuinely empty note set — a failed restore must never read as an
// intact empty context in audit records.
func LoadFromDiskForContext(contextID string) (int, error) {
	s := getNoteStoreForContext(contextID)
	s.mu.Lock()
	defer s.mu.Unlock()

	if s.persistPath == "" {
		return 0, nil
	}

	data, err := os.ReadFile(s.persistPath)
	if err != nil {
		// The only caller is the resume path, where the predecessor saved notes:
		// an absent expected snapshot is itself a state worth surfacing, not one
		// to fold into "zero notes loaded". Wrap ErrNotExist so callers can
		// distinguish absent-expected from corrupt-expected from fresh-empty.
		return 0, fmt.Errorf("notes restore: %w", err)
	}

	loaded := make(map[string]string)
	if err := json.Unmarshal(data, &loaded); err != nil {
		log.Printf("[notes] Warning: failed to parse %s: %v", s.persistPath, err)
		return 0, fmt.Errorf("notes restore failed: corrupt snapshot at %s: %w", s.persistPath, err)
	}

	count := 0
	for k, v := range loaded {
		if _, exists := s.store[k]; !exists {
			s.store[k] = v
			count++
		}
	}

	if count > 0 {
		log.Printf("[notes] Loaded %d notes from: %s (context=%s)", count, s.persistPath, contextID)
	}
	return count, nil
}

// GetAllNotesForContext returns all notes for a specific context.
func GetAllNotesForContext(contextID string) map[string]string {
	s := getNoteStoreForContext(contextID)
	s.mu.RLock()
	defer s.mu.RUnlock()
	result := make(map[string]string, len(s.store))
	for k, v := range s.store {
		result[k] = v
	}
	return result
}

// CleanupContext removes the note store for a deactivated context.
func CleanupContext(contextID string) {
	noteStoresMu.Lock()
	defer noteStoresMu.Unlock()
	delete(noteStores, contextID)
}
