File size: 11,812 Bytes
2e818da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
from __future__ import annotations

import math
import re

from app.schemas.visual_lesson import ArrayAlgorithmBranch, ArrayAlgorithmSpec, ArrayAlgorithmStep, CompiledArrayAlgorithm


SORT_ALGORITHMS = ("bubble_sort", "insertion_sort", "selection_sort", "merge_sort", "quick_sort")
SEARCH_ALGORITHMS = ("linear_search", "binary_search")


class ArrayAlgorithmCompilationError(ValueError):
    pass


class ArrayAlgorithmCompiler:
    @staticmethod
    def parse_prompt(prompt: str) -> tuple[ArrayAlgorithmSpec, list[str]]:
        text = prompt.lower()
        array_match = re.search(r"\[([^\]]+)\]", prompt)
        illustrative = array_match is None
        search_requested = any(term in text for term in ("search", "find ", "lookup"))
        if array_match:
            try:
                values = [float(item.strip()) for item in array_match.group(1).split(",")]
            except ValueError as exc:
                raise ArrayAlgorithmCompilationError("The array must contain comma-separated numbers") from exc
        elif search_requested:
            values = [4, 12, 19, 27, 42, 58, 73, 91]
        else:
            values = [9, 3, 7, 1, 6, 2, 8, 5]
        if not 2 <= len(values) <= 32 or any(not math.isfinite(value) or abs(value) > 1e6 for value in values):
            raise ArrayAlgorithmCompilationError("Array visualizations support 2–32 finite values between -1,000,000 and 1,000,000")

        target_match = re.search(r"(?:search|find|lookup)(?:\s+for)?\s+(-?\d+(?:\.\d+)?)", text)
        target = float(target_match.group(1)) if target_match else (42.0 if search_requested else 0.0)
        if "bubble" in text:
            primary = "bubble_sort"
        elif "insertion" in text:
            primary = "insertion_sort"
        elif "selection" in text:
            primary = "selection_sort"
        elif "merge" in text:
            primary = "merge_sort"
        elif "quick" in text:
            primary = "quick_sort"
        elif "binary" in text:
            primary = "binary_search"
        elif search_requested:
            primary = "linear_search"
        else:
            primary = "quick_sort"

        if primary in SEARCH_ALGORITHMS or search_requested:
            if primary == "binary_search" and values != sorted(values):
                raise ArrayAlgorithmCompilationError("Binary search requires the supplied array to already be sorted")
            enabled = ["linear_search", "binary_search"] if values == sorted(values) else ["linear_search"]
            has_target = True
        else:
            enabled = list(SORT_ALGORITHMS)
            has_target = False
        return ArrayAlgorithmSpec(
            project_id="",
            prompt=prompt,
            input_values=values,
            primary_algorithm=primary,
            enabled_algorithms=enabled,
            search_target=target,
            has_search_target=has_target,
            input_is_illustrative=illustrative,
            assumptions=[
                "Comparisons use ordinary ascending numeric order.",
                "Duplicate values retain their numeric equality; stability depends on the selected algorithm.",
                "Binary search is available only when the supplied array is already sorted.",
            ],
            created_at=0.0,
        ), []

    @staticmethod
    def _sort_branch(values: list[float], algorithm: str) -> ArrayAlgorithmBranch:
        array = list(values)
        steps = [ArrayAlgorithmStep(step_index=0, values=list(array), operation="initial", description="Start with the input array.")]
        comparisons = 0
        writes = 0
        def record(operation: str, description: str, compared: list[int] | None = None, active: list[int] | None = None, sorted_indices: list[int] | None = None, pivot: int = -1, start: int = -1, end: int = -1) -> None:
            steps.append(ArrayAlgorithmStep(step_index=len(steps), values=list(array), compared_indices=compared or [], active_indices=active or [], sorted_indices=sorted_indices or [], pivot_index=pivot, range_start=start, range_end=end, operation=operation, description=description, comparisons=comparisons, writes=writes))

        if algorithm == "bubble_sort":
            n = len(array)
            for end in range(n - 1, 0, -1):
                swapped = False
                for index in range(end):
                    comparisons += 1
                    record("compare", f"Compare positions {index} and {index + 1}.", [index, index + 1], sorted_indices=list(range(end + 1, n)))
                    if array[index] > array[index + 1]:
                        array[index], array[index + 1] = array[index + 1], array[index]
                        writes += 2
                        swapped = True
                        record("swap", "Swap the out-of-order pair.", active=[index, index + 1], sorted_indices=list(range(end + 1, n)))
                if not swapped:
                    break
        elif algorithm == "insertion_sort":
            for index in range(1, len(array)):
                key = array[index]
                cursor = index - 1
                while cursor >= 0:
                    comparisons += 1
                    record("compare", f"Compare the key with position {cursor}.", [cursor, cursor + 1], active=list(range(index + 1)))
                    if array[cursor] <= key:
                        break
                    array[cursor + 1] = array[cursor]
                    writes += 1
                    record("write", "Shift the larger value one position right.", active=[cursor, cursor + 1])
                    cursor -= 1
                array[cursor + 1] = key
                writes += 1
                record("write", "Insert the key into the sorted prefix.", active=list(range(index + 1)))
        elif algorithm == "selection_sort":
            for index in range(len(array) - 1):
                minimum = index
                for cursor in range(index + 1, len(array)):
                    comparisons += 1
                    record("compare", f"Compare the current minimum with position {cursor}.", [minimum, cursor], sorted_indices=list(range(index)))
                    if array[cursor] < array[minimum]:
                        minimum = cursor
                if minimum != index:
                    array[index], array[minimum] = array[minimum], array[index]
                    writes += 2
                    record("swap", "Move the smallest remaining value into place.", active=[index, minimum], sorted_indices=list(range(index + 1)))
        elif algorithm == "merge_sort":
            def merge_sort(start: int, end: int) -> None:
                nonlocal comparisons, writes
                if end - start <= 1:
                    return
                middle = (start + end) // 2
                merge_sort(start, middle)
                merge_sort(middle, end)
                left, right = array[start:middle], array[middle:end]
                li = ri = 0
                merged: list[float] = []
                while li < len(left) and ri < len(right):
                    comparisons += 1
                    record("compare", "Compare the next values from both sorted runs.", [start + li, middle + ri], start=start, end=end - 1)
                    if left[li] <= right[ri]: merged.append(left[li]); li += 1
                    else: merged.append(right[ri]); ri += 1
                merged.extend(left[li:]); merged.extend(right[ri:])
                for offset, value in enumerate(merged):
                    array[start + offset] = value
                    writes += 1
                    record("write", "Write the next merged value.", active=[start + offset], start=start, end=end - 1)
            merge_sort(0, len(array))
        elif algorithm == "quick_sort":
            def quick_sort(low: int, high: int) -> None:
                nonlocal comparisons, writes
                if low >= high:
                    return
                pivot_value = array[high]
                boundary = low
                record("partition", "Use the final value as this partition's pivot.", pivot=high, start=low, end=high)
                for cursor in range(low, high):
                    comparisons += 1
                    record("compare", "Compare the current value with the pivot.", [cursor, high], pivot=high, start=low, end=high)
                    if array[cursor] <= pivot_value:
                        if cursor != boundary:
                            array[cursor], array[boundary] = array[boundary], array[cursor]
                            writes += 2
                            record("swap", "Move the value into the lower partition.", active=[cursor, boundary], pivot=high, start=low, end=high)
                        boundary += 1
                array[boundary], array[high] = array[high], array[boundary]
                writes += 2
                record("swap", "Place the pivot between both partitions.", active=[boundary, high], pivot=boundary, start=low, end=high)
                quick_sort(low, boundary - 1); quick_sort(boundary + 1, high)
            quick_sort(0, len(array) - 1)
        else:
            raise ArrayAlgorithmCompilationError(f"Unsupported sorting algorithm: {algorithm}")
        record("complete", "The array is sorted.", sorted_indices=list(range(len(array))))
        return ArrayAlgorithmBranch(algorithm=algorithm, steps=steps, final_values=list(array), comparisons=comparisons, writes=writes)

    @staticmethod
    def _search_branch(values: list[float], algorithm: str, target: float) -> ArrayAlgorithmBranch:
        steps = [ArrayAlgorithmStep(step_index=0, values=list(values), operation="initial", description=f"Search for {target:g}.")]
        comparisons = 0
        found = -1
        if algorithm == "linear_search":
            for index, value in enumerate(values):
                comparisons += 1
                steps.append(ArrayAlgorithmStep(step_index=len(steps), values=list(values), compared_indices=[index], active_indices=[index], operation="compare", description=f"Compare position {index} with the target.", comparisons=comparisons))
                if value == target:
                    found = index
                    break
        else:
            low, high = 0, len(values) - 1
            while low <= high:
                middle = (low + high) // 2
                comparisons += 1
                steps.append(ArrayAlgorithmStep(step_index=len(steps), values=list(values), compared_indices=[middle], active_indices=list(range(low, high + 1)), range_start=low, range_end=high, operation="compare", description=f"Inspect midpoint {middle} of the remaining range.", comparisons=comparisons))
                if values[middle] == target:
                    found = middle; break
                if values[middle] < target: low = middle + 1
                else: high = middle - 1
        steps.append(ArrayAlgorithmStep(step_index=len(steps), values=list(values), found_index=found, active_indices=[found] if found >= 0 else [], operation="found" if found >= 0 else "not_found", description=f"Target found at index {found}." if found >= 0 else "The target is not present.", comparisons=comparisons))
        return ArrayAlgorithmBranch(algorithm=algorithm, steps=steps, final_values=list(values), found_index=found, comparisons=comparisons)

    def compile_spec(self, spec: ArrayAlgorithmSpec) -> CompiledArrayAlgorithm:
        branches = []
        for algorithm in spec.enabled_algorithms:
            branch = self._search_branch(spec.input_values, algorithm, spec.search_target) if algorithm in SEARCH_ALGORITHMS else self._sort_branch(spec.input_values, algorithm)
            branches.append(branch)
        return CompiledArrayAlgorithm(branches=branches, assertions_passed=True)