Spaces:
Running
Running
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)
|