Skip to content

Commit 2d925d6

Browse files
committed
Preserve entered correction characters
1 parent b11b5e5 commit 2d925d6

3 files changed

Lines changed: 80 additions & 5 deletions

File tree

‎src/codex32/correction.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -914,6 +914,7 @@ def _correct_complete(
914914
primary=frozenset(c.expected_length for c in contexts if c.expected_length is not None),
915915
deadline=deadline,
916916
capture_layers=capture_layers,
917+
observed_text=damaged_text,
917918
)
918919
if result or not complete:
919920
return result, complete

‎src/codex32/indel.py‎

Lines changed: 46 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -284,6 +284,7 @@ def _capacities(erasures: int, _degree: int) -> range:
284284
class _Target:
285285
context: CorrectionContext
286286
text: str
287+
observed_text: str
287288
immutable: int
288289
target: int
289290
base: int
@@ -292,12 +293,18 @@ class _Target:
292293

293294

294295
def _prepare(
295-
context: CorrectionContext, damaged_text: str, classes: Sequence[_StructuralClass]
296+
context: CorrectionContext,
297+
damaged_text: str,
298+
classes: Sequence[_StructuralClass],
299+
observed_text: str | None = None,
296300
) -> _Target | None:
297301
normalized = _normalize(context, damaged_text)
298302
if normalized is None:
299303
return None
300304
text, immutable = normalized
305+
observed = text if observed_text is None else observed_text.replace(" ", "")
306+
if len(observed) != len(text):
307+
raise ValueError("observed text must preserve the searched text length")
301308
target = context.expected_length
302309
assert target is not None
303310
shapes = tuple(shape for shape in classes if shape.delta == len(text) - target)
@@ -311,7 +318,35 @@ def _prepare(
311318
}
312319
base = len(context.hrp) + 1
313320
degree = _checksum_for_encoded_length(context.hrp, target - base).length
314-
return _Target(context, text, immutable, target, base, degree, counts)
321+
return _Target(context, text, observed, immutable, target, base, degree, counts)
322+
323+
324+
def _restore_observed(
325+
candidate: CorrectionCandidate,
326+
state: _Target,
327+
view: _View | None = None,
328+
) -> CorrectionCandidate:
329+
"""Restore diagnostic characters transformed only to make mixed-case text searchable."""
330+
331+
def source_position(position: int) -> int | None:
332+
if view is None:
333+
return position
334+
offset = 0
335+
for start, size in view.spans:
336+
if position < offset + size:
337+
return None if start < 0 else start + position - offset
338+
offset += size
339+
return None
340+
341+
restored = []
342+
body_length = state.target - state.base
343+
for edit in candidate.edits:
344+
position = body_length - edit.reverse_index - 1
345+
source = source_position(position)
346+
if edit.observed and source is not None and 0 <= source < len(state.observed_text) - state.base:
347+
edit = replace(edit, observed=state.observed_text[state.base + source])
348+
restored.append(edit)
349+
return replace(candidate, edits=tuple(restored))
315350

316351

317352
def _layers(
@@ -443,6 +478,8 @@ def _search_fixed(
443478
suspected_profile=state.context.hrp,
444479
immutable_prefix=state.context.immutable_prefix,
445480
)
481+
if fixed is not None:
482+
fixed = _restore_observed(fixed, state)
446483
if fixed is None or not _allowed(state.context, fixed) or allowed is not None and not allowed(fixed):
447484
return None
448485
substitutions = sum(edit.kind == "substitution" for edit in fixed.edits)
@@ -508,10 +545,12 @@ def _search_target(
508545
erasures = tuple(sorted(len(view) - p - 1 for p, _ in unknown))
509546
fixed = solver.correct(
510547
view,
511-
tuple((p, text[state.base + source]) for p, source in unknown if source >= 0),
548+
tuple((p, state.observed_text[state.base + source]) for p, source in unknown if source >= 0),
512549
erasures,
513550
incremental.packed(view),
514551
)
552+
if fixed is not None:
553+
fixed = _restore_observed(fixed, state, view)
515554
if fixed is None or not _allowed(context, fixed) or allowed is not None and not allowed(fixed):
516555
continue
517556
substitutions = sum(edit.kind == "substitution" for edit in fixed.edits)
@@ -520,7 +559,7 @@ def _search_target(
520559
continue
521560
candidate = _adapt(
522561
fixed,
523-
_view_variant(view, text, state.base),
562+
_view_variant(view, state.observed_text, state.base),
524563
state.counts[shape][remaining],
525564
len(text),
526565
state.target,
@@ -541,7 +580,7 @@ def _search_target(
541580
CorrectionEdit(
542581
"transposition",
543582
len(view) - offset - i - 1,
544-
text[state.base + observed_position],
583+
state.observed_text[state.base + observed_position],
545584
candidate.artifact.text[state.base + offset + i],
546585
)
547586
)
@@ -563,6 +602,7 @@ def _search_many(
563602
competitors: bool = False,
564603
allowed: Callable[[CorrectionCandidate], bool] | None = None,
565604
capture_layers: list[tuple[int, int]] | None = None,
605+
observed_text: str | None = None,
566606
) -> tuple[tuple[CorrectionCandidate, ...], bool]:
567607
deadline = monotonic() + 10 if deadline is None else deadline
568608
states = tuple(
@@ -573,6 +613,7 @@ def _search_many(
573613
context,
574614
damaged_text,
575615
_CLASSES,
616+
observed_text,
576617
)
577618
)
578619
is not None

‎tests/test_correction_bch.py‎

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -278,6 +278,39 @@ def test_public_correction_interprets_mixed_case_by_majority(uppercase: bool) ->
278278
assert result[0].artifact.text == source
279279

280280

281+
@pytest.mark.parametrize("uppercase", (False, True))
282+
@pytest.mark.parametrize(("entered", "kind"), (("P", "substitution"), ("B", "erasure")))
283+
def test_mixed_case_correction_edits_preserve_the_entered_character(
284+
uppercase: bool,
285+
entered: str,
286+
kind: str,
287+
) -> None:
288+
source = VECTOR_1["secret_s"].upper() if uppercase else VECTOR_1["secret_s"]
289+
position = 3
290+
observed = entered.lower() if uppercase else entered
291+
damaged = source[:position] + observed + source[position + 1 :]
292+
293+
result = correct(CorrectionContext(Profile.MS, expected_length=len(source)), damaged)
294+
295+
assert len(result) == 1
296+
assert result[0].artifact.text == source
297+
assert tuple(
298+
(edit.kind, edit.reverse_index, edit.observed, edit.replacement) for edit in result[0].edits
299+
) == ((kind, len(source) - position - 1, observed, source[position]),)
300+
301+
302+
def test_mixed_case_structural_edits_preserve_the_entered_character() -> None:
303+
source = VECTOR_1["secret_s"]
304+
damaged = source[:3] + "P" + source[4:20] + source[21:]
305+
306+
result = correct(CorrectionContext(Profile.MS, expected_length=len(source)), damaged)
307+
308+
assert len(result) == 1
309+
assert result[0].artifact.text == source
310+
assert {edit.kind for edit in result[0].edits} == {"insertion", "substitution"}
311+
assert next(edit for edit in result[0].edits if edit.kind == "substitution").observed == "P"
312+
313+
281314
def test_fixed_failures_are_fail_closed() -> None:
282315
mixed = "M" + VECTOR_1["secret_s"][1:]
283316
damaged = list(VECTOR_1["secret_s"])

0 commit comments

Comments
 (0)