package main

import (
	"encoding/json"
	"fmt"
	"io/ioutil"
	"math/rand"
	"os"
	"strconv"
	"strings"
	"time"
)

// MemoryCell represents a single memory cell with type information
type MemoryCell struct {
    Value      interface{}
    Type       string // "int", "float", "string", "array", "stack", "heap"
}

// Memory is a map representing the memory cells, now with enhanced features
type Memory struct {
    cells      map[int]MemoryCell
    stacks     map[int][]interface{} // Stacks for each stack ID
    heaps      map[int]map[int]interface{} // Heaps for each heap ID
    nextStackID int
    nextHeapID  int
}

// NewMemory creates a new Memory instance
func NewMemory() *Memory {
    return &Memory{
        cells:      make(map[int]MemoryCell),
        stacks:     make(map[int][]interface{}),
        heaps:      make(map[int]map[int]interface{}),
        nextStackID: 1,
        nextHeapID:  1,
    }
}

// Get retrieves a value from memory
func (m *Memory) Get(address int) interface{} {
    if cell, exists := m.cells[address]; exists {
        return cell.Value
    }
    return 0 // Default value for backward compatibility
}

// Set stores a value in memory with type information
func (m *Memory) Set(address int, value interface{}, typ string) {
    m.cells[address] = MemoryCell{Value: value, Type: typ}
}

// CreateStack creates a new stack and returns its ID
func (m *Memory) CreateStack() int {
    id := m.nextStackID
    m.stacks[id] = make([]interface{}, 0)
    m.nextStackID++
    return id
}

// Push adds an item to a stack
func (m *Memory) Push(stackID int, value interface{}) {
    if stack, exists := m.stacks[stackID]; exists {
        m.stacks[stackID] = append(stack, value)
    }
}

// Pop removes and returns the top item from a stack
func (m *Memory) Pop(stackID int) interface{} {
    if stack, exists := m.stacks[stackID]; exists && len(stack) > 0 {
        top := stack[len(stack)-1]
        m.stacks[stackID] = stack[:len(stack)-1]
        return top
    }
    return 0 // Default value for backward compatibility
}

// CreateHeap creates a new heap and returns its ID
func (m *Memory) CreateHeap() int {
    id := m.nextHeapID
    m.heaps[id] = make(map[int]interface{})
    m.nextHeapID++
    return id
}

// HeapSet stores a value in a heap
func (m *Memory) HeapSet(heapID, address int, value interface{}) {
    if heap, exists := m.heaps[heapID]; exists {
        heap[address] = value
    }
}

// HeapGet retrieves a value from a heap
func (m *Memory) HeapGet(heapID, address int) interface{} {
    if heap, exists := m.heaps[heapID]; exists {
        if val, exists := heap[address]; exists {
            return val
        }
    }
    return 0 // Default value for backward compatibility
}

// SVM is the Sapbot VM interpreter
type SVM struct {
	pc   int
	mem  Memory
	code map[int]string
}

// NewSVM creates a new SVM instance
func NewSVM() *SVM {
	return &SVM{
		pc:   0,
		mem:  *NewMemory(),
		code: make(map[int]string),
	}
}

// LoadProgram loads a program from a file (JSON or SVM2 format)
func (svm *SVM) LoadProgram(filename string) error {
	file, err := ioutil.ReadFile(filename)
	if err != nil {
		return err
	}

	// Check file extension to determine format
	if strings.HasSuffix(filename, ".svm2") {
		// Parse SVM2 format
		lines := strings.Split(string(file), "\n")
		for _, line := range lines {
			// Skip empty lines and comments
			line = strings.TrimSpace(line)
			if line == "" || strings.HasPrefix(line, "#") {
				continue
			}

			// Parse line number and instruction
			parts := strings.Fields(line)
			if len(parts) < 2 {
				continue // Skip malformed lines
			}

			lineNum, err := strconv.Atoi(parts[0])
			if err != nil {
				return fmt.Errorf("invalid line number: %s", parts[0])
			}

			// The rest is the instruction
			instruction := strings.Join(parts[1:], " ")
			svm.code[lineNum] = instruction
		}
	} else {
		// Parse JSON format (original behavior)
		// Try to unmarshal as a map (object)
		var rawMap map[string]interface{}
		if err := json.Unmarshal(file, &rawMap); err == nil {
			// Parse the JSON object into the code map
			for key, value := range rawMap {
				lineNum, err := strconv.Atoi(key)
				if err != nil {
					return fmt.Errorf("invalid line number: %s", key)
				}
				svm.code[lineNum] = value.(string)
			}
		} else {
			// Try to unmarshal as a slice (array)
			var rawSlice []string
			if err := json.Unmarshal(file, &rawSlice); err != nil {
				return err
			}
			// For arrays, assume each line is a command and assign sequential line numbers
			for i, cmd := range rawSlice {
				svm.code[i+1] = cmd
			}
		}
	}

	return nil
}

// ParseExpr parses a value expression (e.g., V3, M3, Ttext, F0.5, S1:3, H1:5)
func (svm *SVM) ParseExpr(val string) interface{} {
	if len(val) == 0 {
		return 0
	}

	switch val[0] {
	case 'V':
		// Integer value
		intVal, _ := strconv.Atoi(val[1:])
		return intVal
	case 'M':
		// Memory cell reference
		cell, _ := strconv.Atoi(val[1:])
		return svm.mem.Get(cell)
	case 'T':
		// Text value
		return val[1:]
	case 'F':
		// Float value
		floatVal, _ := strconv.ParseFloat(val[1:], 64)
		return floatVal
	case 'S':
		// Stack reference: S<stackID>:<index>
		parts := strings.Split(val[1:], ":")
		if len(parts) != 2 {
			return 0
		}
		stackID, _ := strconv.Atoi(parts[0])
		index, _ := strconv.Atoi(parts[1])
		stack := svm.mem.stacks[stackID]
		if index < len(stack) {
			return stack[index]
		}
		return 0
	case 'H':
		// Heap reference: H<heapID>:<address>
		parts := strings.Split(val[1:], ":")
		if len(parts) != 2 {
			return 0
		}
		heapID, _ := strconv.Atoi(parts[0])
		address, _ := strconv.Atoi(parts[1])
		return svm.mem.HeapGet(heapID, address)
	case 'A':
		// Array reference: A<baseAddress>:<index>
		parts := strings.Split(val[1:], ":")
		if len(parts) != 2 {
			return 0
		}
		baseAddress, _ := strconv.Atoi(parts[0])
		index, _ := strconv.Atoi(parts[1])
		// For simplicity, treat arrays as sequential memory cells
		// In a real implementation, you might want a more sophisticated array type
		return svm.mem.Get(baseAddress + index)
	default:
		return 0
	}
}

// RunError represents an error that occurred during program execution
type RunError struct {
	LineNumber int
	Message    string
}

func (e *RunError) Error() string {
	return fmt.Sprintf("Error at line %d: %s", e.LineNumber, e.Message)
}

// Run executes the loaded program
func (svm *SVM) Run() error {
	rand.Seed(time.Now().UnixNano())

	for svm.pc < 65000 {
		line, exists := svm.code[svm.pc]
		if !exists {
			svm.pc++
			continue
		}

		parts := strings.Split(line, ";")
		if len(parts) == 0 {
			svm.pc++
			continue
		}

		instruction := parts[0]
		args := parts[1:]

		var err error
		switch instruction {
			// Basic instructions
			case "LOAD":
				if len(args) != 2 {
					return &RunError{svm.pc, "Invalid number of operands for LOAD"}
				}
				cell := svm.ParseExpr(args[0]).(int)
				value := svm.ParseExpr(args[1])
				svm.mem.Set(cell, value, "int") // Default type for backward compatibility

			case "NOP":
				// Do nothing

			case "GOTO":
				if len(args) != 1 {
					return &RunError{svm.pc, "Invalid number of operands for GOTO"}
				}
				svm.pc = svm.ParseExpr(args[0]).(int)
				continue

			case "DEBUG":
				if len(args) != 1 {
					return &RunError{svm.pc, "Invalid number of operands for DEBUG"}
				}
				fmt.Println(svm.ParseExpr(args[0]))

			case "NOT":
				if len(args) != 2 {
					return &RunError{svm.pc, "Invalid number of operands for NOT"}
				}
				cell := svm.ParseExpr(args[1]).(int)
				value := svm.ParseExpr(args[0])
				if boolVal, ok := value.(int); ok {
					svm.mem.Set(cell, 0, "int")
					if boolVal != 0 {
						svm.mem.Set(cell, 1, "int")
					}
				} else {
					svm.mem.Set(cell, 0, "int")
				}

			case "ADD":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for ADD"}
				}
				a := svm.ParseExpr(args[0])
				b := svm.ParseExpr(args[1])
				to := svm.ParseExpr(args[2]).(int)

				var result interface{}
				switch a := a.(type) {
				case int:
					bVal := b.(int)
					result = a + bVal
				case float64:
					bVal := b.(float64)
					result = a + bVal
				default:
					result = 0
				}
				svm.mem.Set(to, result, "int")

			case "SUB":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for SUB"}
				}
				a := svm.ParseExpr(args[0])
				b := svm.ParseExpr(args[1])
				to := svm.ParseExpr(args[2]).(int)

				var result interface{}
				switch a := a.(type) {
				case int:
					bVal := b.(int)
					result = a - bVal
				case float64:
					bVal := b.(float64)
					result = a - bVal
				default:
					result = 0
				}
				svm.mem.Set(to, result, "int")

			case "DIV":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for DIV"}
				}
				a := svm.ParseExpr(args[0])
				b := svm.ParseExpr(args[1])
				to := svm.ParseExpr(args[2]).(int)

				var result interface{}
				switch a := a.(type) {
				case int:
					bVal := b.(int)
					if bVal == 0 {
						result = 0
					} else {
						result = a / bVal
					}
				case float64:
					bVal := b.(float64)
					if bVal == 0 {
						result = 0.0
					} else {
						result = a / bVal
					}
				default:
					result = 0
				}
				svm.mem.Set(to, result, "int")

			case "MUL":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for MUL"}
				}
				a := svm.ParseExpr(args[0])
				b := svm.ParseExpr(args[1])
				to := svm.ParseExpr(args[2]).(int)

				var result interface{}
				switch a := a.(type) {
				case int:
					bVal := b.(int)
					result = a * bVal
				case float64:
					bVal := b.(float64)
					result = a * bVal
				default:
					result = 0
				}
				svm.mem.Set(to, result, "int")

			case "MOD":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for MOD"}
				}
				a := svm.ParseExpr(args[0]).(int)
				b := svm.ParseExpr(args[1]).(int)
				to := svm.ParseExpr(args[2]).(int)

				if b == 0 {
					svm.mem.Set(to, 0, "int")
				} else {
					svm.mem.Set(to, a%b, "int")
				}

			case "ALB":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for ALB"}
				}
				a := svm.ParseExpr(args[0])
				b := svm.ParseExpr(args[1])
				to := svm.ParseExpr(args[2]).(int)

				var aVal, bVal int
				switch a := a.(type) {
				case int:
					aVal = a
				case float64:
					aVal = int(a)
				default:
					aVal = 0
				}

				switch b := b.(type) {
				case int:
					bVal = b
				case float64:
					bVal = int(b)
				default:
					bVal = 0
				}

				if aVal < bVal {
					svm.pc = to
					continue
				}

			case "AQB":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for AQB"}
				}
				a := svm.ParseExpr(args[0])
				b := svm.ParseExpr(args[1])
				to := svm.ParseExpr(args[2]).(int)

				var aVal, bVal int
				switch a := a.(type) {
				case int:
					aVal = a
				case float64:
					aVal = int(a)
				default:
					aVal = 0
				}

				switch b := b.(type) {
				case int:
					bVal = b
				case float64:
					bVal = int(b)
				default:
					bVal = 0
				}

				if aVal == bVal {
					svm.pc = to
					continue
				}

			case "ABB":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for ABB"}
				}
				a := svm.ParseExpr(args[0])
				b := svm.ParseExpr(args[1])
				to := svm.ParseExpr(args[2]).(int)

				var aVal, bVal int
				switch a := a.(type) {
				case int:
					aVal = a
				case float64:
					aVal = int(a)
				default:
					aVal = 0
				}

				switch b := b.(type) {
				case int:
					bVal = b
				case float64:
					bVal = int(b)
				default:
					bVal = 0
				}

				if aVal > bVal {
					svm.pc = to
					continue
				}

			case "AND":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for AND"}
				}
				a := svm.ParseExpr(args[0])
				b := svm.ParseExpr(args[1])
				to := svm.ParseExpr(args[2]).(int)

				var aVal, bVal int
				switch a := a.(type) {
				case int:
					aVal = a
				case float64:
					aVal = int(a)
				default:
					aVal = 0
				}

				switch b := b.(type) {
				case int:
					bVal = b
				case float64:
					bVal = int(b)
				default:
					bVal = 0
				}

				if aVal != 0 && bVal != 0 {
					svm.mem.Set(to, 1, "int")
				} else {
					svm.mem.Set(to, 0, "int")
				}

			case "OR":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for OR"}
				}
				a := svm.ParseExpr(args[0])
				b := svm.ParseExpr(args[1])
				to := svm.ParseExpr(args[2]).(int)

				var aVal, bVal int
				switch a := a.(type) {
				case int:
					aVal = a
				case float64:
					aVal = int(a)
				default:
					aVal = 0
				}

				switch b := b.(type) {
				case int:
					bVal = b
				case float64:
					bVal = int(b)
				default:
					bVal = 0
				}

				if aVal != 0 || bVal != 0 {
					svm.mem.Set(to, 1, "int")
				} else {
					svm.mem.Set(to, 0, "int")
				}

			case "DEBINP":
				if len(args) != 1 {
					return &RunError{svm.pc, "Invalid number of operands for DEBINP"}
				}
				cell := svm.ParseExpr(args[0]).(int)
				var input string
				fmt.Print("Input: ")
				fmt.Scanln(&input)
				svm.mem.Set(cell, input, "string")

			case "PARSEINT":
				if len(args) != 2 {
					return &RunError{svm.pc, "Invalid number of operands for PARSEINT"}
				}
				in := svm.ParseExpr(args[0]).(string)
				out := svm.ParseExpr(args[1]).(int)

				intVal, err := strconv.Atoi(in)
				if err != nil {
					svm.mem.Set(out, 0, "int")
				} else {
					svm.mem.Set(out, intVal, "int")
				}

			case "RND":
				if len(args) != 2 {
					return &RunError{svm.pc, "Invalid number of operands for RND"}
				}
				max := svm.ParseExpr(args[0]).(int)
				to := svm.ParseExpr(args[1]).(int)
				svm.mem.Set(to, rand.Intn(max), "int")

			// Stack operations
			case "STACK_CREATE":
				if len(args) != 1 {
					return &RunError{svm.pc, "Invalid number of operands for STACK_CREATE"}
				}
				to := svm.ParseExpr(args[0]).(int)
				stackID := svm.mem.CreateStack()
				svm.mem.Set(to, stackID, "stack")

			case "STACK_PUSH":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for STACK_PUSH"}
				}
				stackID := svm.ParseExpr(args[0]).(int)
				value := svm.ParseExpr(args[1])
				svm.mem.Push(stackID, value)

			case "STACK_POP":
				if len(args) != 2 {
					return &RunError{svm.pc, "Invalid number of operands for STACK_POP"}
				}
				stackID := svm.ParseExpr(args[0]).(int)
				to := svm.ParseExpr(args[1]).(int)
				value := svm.mem.Pop(stackID)
				svm.mem.Set(to, value, "int")

			// Heap operations
			case "HEAP_CREATE":
				if len(args) != 1 {
					return &RunError{svm.pc, "Invalid number of operands for HEAP_CREATE"}
				}
				to := svm.ParseExpr(args[0]).(int)
				heapID := svm.mem.CreateHeap()
				svm.mem.Set(to, heapID, "heap")

			case "HEAP_SET":
				if len(args) != 4 {
					return &RunError{svm.pc, "Invalid number of operands for HEAP_SET"}
				}
				heapID := svm.ParseExpr(args[0]).(int)
				address := svm.ParseExpr(args[1]).(int)
				value := svm.ParseExpr(args[2])
				svm.mem.HeapSet(heapID, address, value)

			case "HEAP_GET":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for HEAP_GET"}
				}
				heapID := svm.ParseExpr(args[0]).(int)
				address := svm.ParseExpr(args[1]).(int)
				to := svm.ParseExpr(args[2]).(int)
				value := svm.mem.HeapGet(heapID, address)
				svm.mem.Set(to, value, "int")

			// Array operations (simplified)
			case "ARRAY_SET":
				if len(args) != 4 {
					return &RunError{svm.pc, "Invalid number of operands for ARRAY_SET"}
				}
				baseAddress := svm.ParseExpr(args[0]).(int)
				index := svm.ParseExpr(args[1]).(int)
				value := svm.ParseExpr(args[2])
				svm.mem.Set(baseAddress+index, value, "int")

			case "ARRAY_GET":
				if len(args) != 3 {
					return &RunError{svm.pc, "Invalid number of operands for ARRAY_GET"}
				}
				baseAddress := svm.ParseExpr(args[0]).(int)
				index := svm.ParseExpr(args[1]).(int)
				to := svm.ParseExpr(args[2]).(int)
				value := svm.mem.Get(baseAddress + index)
				svm.mem.Set(to, value, "int")

			// Function support (basic)
			case "CALL":
				if len(args) != 2 {
					return &RunError{svm.pc, "Invalid number of operands for CALL"}
				}
				// For now, just jump to the address (real function support would require a stack for return addresses)
				address := svm.ParseExpr(args[0]).(int)
				svm.pc = address
				continue

            case "RET":
                // For now, just continue (real function support would require a stack for return addresses)
                continue

            // String manipulation instructions
            case "STRSPLIT":
                if len(args) != 3 {
                    return &RunError{svm.pc, "Invalid number of operands for STRSPLIT"}
                }
                text := svm.ParseExpr(args[0]).(string)
                delimiter := svm.ParseExpr(args[1]).(string)
                baseAddress := svm.ParseExpr(args[2]).(int)

                // Split the text by delimiter
                parts := strings.Split(text, delimiter)

                // Store each part in memory starting at baseAddress
                for i, part := range parts {
                    svm.mem.Set(baseAddress+i, part, "string")
                }

            case "STRCONCAT":
                if len(args) != 3 {
                    return &RunError{svm.pc, "Invalid number of operands for STRCONCAT"}
                }
                baseAddress := svm.ParseExpr(args[0]).(int)
                length := svm.ParseExpr(args[1]).(int)
                to := svm.ParseExpr(args[2]).(int)

                // Concatenate strings from memory
                var result strings.Builder
                for i := 0; i < length; i++ {
                    val := svm.mem.Get(baseAddress + i)
                    if strVal, ok := val.(string); ok {
                        result.WriteString(strVal)
                    }
                }
                svm.mem.Set(to, result.String(), "string")

            case "STRTOASCII":
                if len(args) != 2 {
                    return &RunError{svm.pc, "Invalid number of operands for STRTOASCII"}
                }
                text := svm.ParseExpr(args[0]).(string)
                baseAddress := svm.ParseExpr(args[1]).(int)

                // Convert each character to ASCII code
                for i, char := range text {
                    svm.mem.Set(baseAddress+i, int(char), "int")
                }

            case "ASCIITOSTR":
                if len(args) != 3 {
                    return &RunError{svm.pc, "Invalid number of operands for ASCIITOSTR"}
                }
                baseAddress := svm.ParseExpr(args[0]).(int)
                length := svm.ParseExpr(args[1]).(int)
                to := svm.ParseExpr(args[2]).(int)

                // Convert ASCII codes to string
                var result strings.Builder
                for i := 0; i < length; i++ {
                    asciiCode := svm.mem.Get(baseAddress + i).(int)
                    result.WriteByte(byte(asciiCode))
                }
                svm.mem.Set(to, result.String(), "string")

			default:
				return &RunError{svm.pc, fmt.Sprintf("Unknown instruction: %s", instruction)}
		}

		svm.pc++
	}
}

func main() {
	if len(os.Args) < 2 {
		fmt.Println("Usage: go run svm.go <program.svm>")
		return
	}

	svm := NewSVM()
	if err := svm.LoadProgram(os.Args[1]); err != nil {
		fmt.Printf("Error loading program: %v\n", err)
		return
	}

	if err := svm.Run(); err != nil {
		fmt.Printf("Runtime error: %v\n", err)
		return
	}
}
