Added an interface for validators, and implemented an initial deterministic validator to protect terms in the provided glossary
This commit is contained in:
@@ -40,8 +40,9 @@ class ReportedSkip:
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ReviewStageReport:
|
||||
class ValidatorReport:
|
||||
name: str
|
||||
execution_kind: str
|
||||
candidate_count: int
|
||||
approved_count: int
|
||||
rejected_count: int
|
||||
@@ -50,16 +51,6 @@ class ReviewStageReport:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DeterministicFilterReport:
|
||||
name: str
|
||||
passed_count: int
|
||||
rejected_count: int
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModuleRunReport:
|
||||
instance_name: str
|
||||
@@ -67,8 +58,7 @@ class ModuleRunReport:
|
||||
replacement_policy: str
|
||||
section_count: int
|
||||
proposal_count: int
|
||||
deterministic_filters: List[DeterministicFilterReport]
|
||||
review_stages: List[ReviewStageReport]
|
||||
validators: List[ValidatorReport]
|
||||
approved_count: int
|
||||
applied_count: int
|
||||
skipped_count: int
|
||||
@@ -80,8 +70,7 @@ class ModuleRunReport:
|
||||
"replacement_policy": self.replacement_policy,
|
||||
"section_count": self.section_count,
|
||||
"proposal_count": self.proposal_count,
|
||||
"deterministic_filters": [item.to_dict() for item in self.deterministic_filters],
|
||||
"review_stages": [item.to_dict() for item in self.review_stages],
|
||||
"validators": [item.to_dict() for item in self.validators],
|
||||
"approved_count": self.approved_count,
|
||||
"applied_count": self.applied_count,
|
||||
"skipped_count": self.skipped_count,
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, List, Optional, Protocol, Sequence
|
||||
from typing import Any, Protocol, Sequence
|
||||
|
||||
from audita.core.config import AuditaConfig
|
||||
from audita.core.schemas import Glossary, TranscriptSegment
|
||||
from audita.validators.base import Validator
|
||||
|
||||
|
||||
ReplacementPolicy = str
|
||||
@@ -36,20 +37,6 @@ class CorrectionProposal:
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FilterDecision:
|
||||
approved: bool
|
||||
reason: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ReviewDecision:
|
||||
proposal_index: int
|
||||
approved: bool
|
||||
confidence: Optional[float] = None
|
||||
reason: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModuleContext:
|
||||
run_spec: ModuleRunSpec
|
||||
@@ -70,42 +57,11 @@ class StructuredLLMClient(Protocol):
|
||||
...
|
||||
|
||||
|
||||
class DeterministicFilter(Protocol):
|
||||
name: str
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
proposal: CorrectionProposal,
|
||||
transcript: Sequence[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
) -> FilterDecision:
|
||||
...
|
||||
|
||||
|
||||
class ReviewStage(Protocol):
|
||||
name: str
|
||||
|
||||
def review(
|
||||
self,
|
||||
proposals: Sequence[CorrectionProposal],
|
||||
transcript: Sequence[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
llm_client: Optional[StructuredLLMClient],
|
||||
run_dir: Path,
|
||||
) -> Sequence[ReviewDecision]:
|
||||
...
|
||||
|
||||
|
||||
class TranscriptModule(Protocol):
|
||||
module_key: str
|
||||
replacement_policy: ReplacementPolicy
|
||||
|
||||
def deterministic_filters(self) -> Sequence[DeterministicFilter]:
|
||||
...
|
||||
|
||||
def review_stages(self) -> Sequence[ReviewStage]:
|
||||
def validators(self) -> Sequence[Validator]:
|
||||
...
|
||||
|
||||
def propose(
|
||||
|
||||
@@ -4,23 +4,11 @@ from typing import Callable, Dict, List, Optional, Sequence, Tuple, Union
|
||||
|
||||
from audita.core.chunking import TranscriptSection, chunk_transcript
|
||||
from audita.core.config import AuditaConfig
|
||||
from audita.core.reporting import (
|
||||
AppliedChange,
|
||||
DeterministicFilterReport,
|
||||
ModuleRunReport,
|
||||
ReportedSkip,
|
||||
ReviewStageReport,
|
||||
)
|
||||
from audita.core.reporting import AppliedChange, ModuleRunReport, ReportedSkip, ValidatorReport
|
||||
from audita.core.schemas import Glossary, TranscriptSegment
|
||||
from audita.validators.base import ValidationContext, ValidationDecision, ValidationResult
|
||||
|
||||
from .models import (
|
||||
CorrectionProposal,
|
||||
FilterDecision,
|
||||
ModuleContext,
|
||||
ModuleRunSpec,
|
||||
ReviewDecision,
|
||||
StructuredLLMClient,
|
||||
)
|
||||
from .models import CorrectionProposal, ModuleContext, ModuleRunSpec, StructuredLLMClient
|
||||
|
||||
|
||||
ProgressCallback = Callable[[str], None]
|
||||
@@ -115,47 +103,15 @@ def _run_module(
|
||||
]
|
||||
|
||||
surviving = proposals
|
||||
filter_reports: List[DeterministicFilterReport] = []
|
||||
validator_reports: List[ValidatorReport] = []
|
||||
skipped: List[ReportedSkip] = []
|
||||
for deterministic_filter in module.deterministic_filters():
|
||||
next_survivors: List[CorrectionProposal] = []
|
||||
rejected_count = 0
|
||||
for proposal in surviving:
|
||||
decision = deterministic_filter.evaluate(proposal, working, context.glossary, context.config)
|
||||
if decision.approved:
|
||||
next_survivors.append(proposal)
|
||||
continue
|
||||
rejected_count += 1
|
||||
skipped.append(
|
||||
ReportedSkip(
|
||||
module_instance=proposal.module_instance,
|
||||
module_key=proposal.module_key,
|
||||
proposal_index=proposal.proposal_index,
|
||||
id=proposal.id,
|
||||
reason=decision.reason or f"{deterministic_filter.name} rejected proposal",
|
||||
original_text=proposal.original_text,
|
||||
corrected_text=proposal.corrected_text,
|
||||
confidence=proposal.confidence,
|
||||
actual_text=_segment_text_by_id(working).get(proposal.id),
|
||||
source=f"deterministic_filter:{deterministic_filter.name}",
|
||||
)
|
||||
)
|
||||
filter_reports.append(
|
||||
DeterministicFilterReport(
|
||||
name=deterministic_filter.name,
|
||||
passed_count=len(next_survivors),
|
||||
rejected_count=rejected_count,
|
||||
)
|
||||
)
|
||||
surviving = next_survivors
|
||||
|
||||
review_reports: List[ReviewStageReport] = []
|
||||
for review_stage in module.review_stages():
|
||||
for validator in module.validators():
|
||||
candidate_count = len(surviving)
|
||||
if not surviving:
|
||||
review_reports.append(
|
||||
ReviewStageReport(
|
||||
name=review_stage.name,
|
||||
validator_reports.append(
|
||||
ValidatorReport(
|
||||
name=validator.name,
|
||||
execution_kind=validator.execution_kind,
|
||||
candidate_count=0,
|
||||
approved_count=0,
|
||||
rejected_count=0,
|
||||
@@ -163,8 +119,17 @@ def _run_module(
|
||||
)
|
||||
continue
|
||||
|
||||
decisions = list(review_stage.review(surviving, working, context.glossary, context.config, llm_client, context.run_dir))
|
||||
decisions_by_index = _index_review_decisions(decisions, surviving, review_stage.name)
|
||||
validation_context = ValidationContext(
|
||||
proposals=surviving,
|
||||
transcript=working,
|
||||
glossary=context.glossary,
|
||||
config=context.config,
|
||||
run_spec=context.run_spec,
|
||||
run_dir=context.run_dir,
|
||||
llm_client=llm_client,
|
||||
)
|
||||
result = validator.validate(validation_context)
|
||||
decisions_by_index = _index_validation_decisions(result, surviving, validator.name)
|
||||
approved: List[CorrectionProposal] = []
|
||||
rejected_count = 0
|
||||
for proposal in surviving:
|
||||
@@ -179,17 +144,18 @@ def _run_module(
|
||||
module_key=proposal.module_key,
|
||||
proposal_index=proposal.proposal_index,
|
||||
id=proposal.id,
|
||||
reason=decision.reason or f"{review_stage.name} rejected proposal",
|
||||
reason=decision.reason or f"{validator.name} rejected proposal",
|
||||
original_text=proposal.original_text,
|
||||
corrected_text=proposal.corrected_text,
|
||||
confidence=proposal.confidence,
|
||||
actual_text=_segment_text_by_id(working).get(proposal.id),
|
||||
source=f"review_stage:{review_stage.name}",
|
||||
source=f"validator:{validator.name}",
|
||||
)
|
||||
)
|
||||
review_reports.append(
|
||||
ReviewStageReport(
|
||||
name=review_stage.name,
|
||||
validator_reports.append(
|
||||
ValidatorReport(
|
||||
name=validator.name,
|
||||
execution_kind=validator.execution_kind,
|
||||
candidate_count=candidate_count,
|
||||
approved_count=len(approved),
|
||||
rejected_count=rejected_count,
|
||||
@@ -213,8 +179,7 @@ def _run_module(
|
||||
replacement_policy=module.replacement_policy,
|
||||
section_count=len(sections),
|
||||
proposal_count=len(proposals),
|
||||
deterministic_filters=filter_reports,
|
||||
review_stages=review_reports,
|
||||
validators=validator_reports,
|
||||
approved_count=len(surviving),
|
||||
applied_count=len(applied_changes),
|
||||
skipped_count=len(skipped),
|
||||
@@ -227,22 +192,22 @@ def _run_module(
|
||||
)
|
||||
|
||||
|
||||
def _index_review_decisions(
|
||||
decisions: Sequence[ReviewDecision],
|
||||
def _index_validation_decisions(
|
||||
result: ValidationResult,
|
||||
proposals: Sequence[CorrectionProposal],
|
||||
stage_name: str,
|
||||
) -> Dict[int, ReviewDecision]:
|
||||
validator_name: str,
|
||||
) -> Dict[int, ValidationDecision]:
|
||||
expected_indexes = {proposal.proposal_index for proposal in proposals}
|
||||
indexed: Dict[int, ReviewDecision] = {}
|
||||
for decision in decisions:
|
||||
indexed: Dict[int, ValidationDecision] = {}
|
||||
for decision in result.decisions:
|
||||
if decision.proposal_index in indexed:
|
||||
raise ValueError(f"Review stage '{stage_name}' returned duplicate proposal indexes.")
|
||||
raise ValueError(f"Validator '{validator_name}' returned duplicate proposal indexes.")
|
||||
if decision.proposal_index not in expected_indexes:
|
||||
raise ValueError(f"Review stage '{stage_name}' returned an unknown proposal index.")
|
||||
raise ValueError(f"Validator '{validator_name}' returned an unknown proposal index.")
|
||||
indexed[decision.proposal_index] = decision
|
||||
missing = expected_indexes - set(indexed)
|
||||
if missing:
|
||||
raise ValueError(f"Review stage '{stage_name}' omitted proposal indexes: {sorted(missing)}")
|
||||
raise ValueError(f"Validator '{validator_name}' omitted proposal indexes: {sorted(missing)}")
|
||||
return indexed
|
||||
|
||||
|
||||
|
||||
@@ -1,61 +1,19 @@
|
||||
from typing import Sequence
|
||||
|
||||
from audita.core.config import AuditaConfig
|
||||
from audita.core.schemas import Glossary, TranscriptSegment
|
||||
from audita.framework.models import (
|
||||
CorrectionProposal,
|
||||
DeterministicFilter,
|
||||
FilterDecision,
|
||||
ModuleContext,
|
||||
ReviewDecision,
|
||||
ReviewStage,
|
||||
)
|
||||
|
||||
|
||||
class _StubFilter:
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
proposal: CorrectionProposal,
|
||||
transcript: Sequence[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
) -> FilterDecision:
|
||||
return FilterDecision(approved=True)
|
||||
|
||||
|
||||
class _StubReviewStage:
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
|
||||
def review(
|
||||
self,
|
||||
proposals: Sequence[CorrectionProposal],
|
||||
transcript: Sequence[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
llm_client,
|
||||
run_dir,
|
||||
) -> Sequence[ReviewDecision]:
|
||||
return [ReviewDecision(proposal_index=proposal.proposal_index, approved=True) for proposal in proposals]
|
||||
from audita.core.schemas import TranscriptSegment
|
||||
from audita.framework.models import CorrectionProposal, ModuleContext
|
||||
from audita.validators import ProtectedGlossaryTermsValidator, StubLLMValidator, Validator
|
||||
|
||||
|
||||
class GlossaryModule:
|
||||
module_key = "glossary"
|
||||
replacement_policy = "replace_all"
|
||||
|
||||
def deterministic_filters(self) -> Sequence[DeterministicFilter]:
|
||||
def validators(self) -> Sequence[Validator]:
|
||||
return [
|
||||
_StubFilter("glossary_direction_guard"),
|
||||
_StubFilter("protected_glossary_guard"),
|
||||
]
|
||||
|
||||
def review_stages(self) -> Sequence[ReviewStage]:
|
||||
return [
|
||||
_StubReviewStage("toward_glossary_term_review"),
|
||||
_StubReviewStage("context_support_review"),
|
||||
ProtectedGlossaryTermsValidator("protected_glossary_guard"),
|
||||
StubLLMValidator("toward_glossary_term_review"),
|
||||
StubLLMValidator("context_support_review"),
|
||||
]
|
||||
|
||||
def propose(
|
||||
|
||||
@@ -1,60 +1,18 @@
|
||||
from typing import Sequence
|
||||
|
||||
from audita.core.config import AuditaConfig
|
||||
from audita.core.schemas import Glossary, TranscriptSegment
|
||||
from audita.framework.models import (
|
||||
CorrectionProposal,
|
||||
DeterministicFilter,
|
||||
FilterDecision,
|
||||
ModuleContext,
|
||||
ReviewDecision,
|
||||
ReviewStage,
|
||||
)
|
||||
|
||||
|
||||
class _StubFilter:
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
proposal: CorrectionProposal,
|
||||
transcript: Sequence[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
) -> FilterDecision:
|
||||
return FilterDecision(approved=True)
|
||||
|
||||
|
||||
class _StubReviewStage:
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
|
||||
def review(
|
||||
self,
|
||||
proposals: Sequence[CorrectionProposal],
|
||||
transcript: Sequence[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
llm_client,
|
||||
run_dir,
|
||||
) -> Sequence[ReviewDecision]:
|
||||
return [ReviewDecision(proposal_index=proposal.proposal_index, approved=True) for proposal in proposals]
|
||||
from audita.core.schemas import TranscriptSegment
|
||||
from audita.framework.models import CorrectionProposal, ModuleContext
|
||||
from audita.validators import ProtectedGlossaryTermsValidator, StubLLMValidator, Validator
|
||||
|
||||
|
||||
class GrammarModule:
|
||||
module_key = "grammar"
|
||||
replacement_policy = "require_unique"
|
||||
|
||||
def deterministic_filters(self) -> Sequence[DeterministicFilter]:
|
||||
def validators(self) -> Sequence[Validator]:
|
||||
return [
|
||||
_StubFilter("protected_glossary_guard"),
|
||||
_StubFilter("punctuation_capitalization_spacing_guard"),
|
||||
]
|
||||
|
||||
def review_stages(self) -> Sequence[ReviewStage]:
|
||||
return [
|
||||
_StubReviewStage("edited_text_readability_review"),
|
||||
ProtectedGlossaryTermsValidator("protected_glossary_guard"),
|
||||
StubLLMValidator("edited_text_readability_review"),
|
||||
]
|
||||
|
||||
def propose(
|
||||
|
||||
@@ -1,62 +1,20 @@
|
||||
from typing import Sequence
|
||||
|
||||
from audita.core.config import AuditaConfig
|
||||
from audita.core.schemas import Glossary, TranscriptSegment
|
||||
from audita.framework.models import (
|
||||
CorrectionProposal,
|
||||
DeterministicFilter,
|
||||
FilterDecision,
|
||||
ModuleContext,
|
||||
ReviewDecision,
|
||||
ReviewStage,
|
||||
)
|
||||
|
||||
|
||||
class _StubFilter:
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
proposal: CorrectionProposal,
|
||||
transcript: Sequence[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
) -> FilterDecision:
|
||||
return FilterDecision(approved=True)
|
||||
|
||||
|
||||
class _StubReviewStage:
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
|
||||
def review(
|
||||
self,
|
||||
proposals: Sequence[CorrectionProposal],
|
||||
transcript: Sequence[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
llm_client,
|
||||
run_dir,
|
||||
) -> Sequence[ReviewDecision]:
|
||||
return [ReviewDecision(proposal_index=proposal.proposal_index, approved=True) for proposal in proposals]
|
||||
from audita.core.schemas import TranscriptSegment
|
||||
from audita.framework.models import CorrectionProposal, ModuleContext
|
||||
from audita.validators import ProtectedGlossaryTermsValidator, StubLLMValidator, Validator
|
||||
|
||||
|
||||
class HomophonesModule:
|
||||
module_key = "homophones"
|
||||
replacement_policy = "require_unique"
|
||||
|
||||
def deterministic_filters(self) -> Sequence[DeterministicFilter]:
|
||||
def validators(self) -> Sequence[Validator]:
|
||||
return [
|
||||
_StubFilter("protected_glossary_guard"),
|
||||
_StubFilter("short_span_guard"),
|
||||
]
|
||||
|
||||
def review_stages(self) -> Sequence[ReviewStage]:
|
||||
return [
|
||||
_StubReviewStage("acoustic_similarity_review"),
|
||||
_StubReviewStage("contextual_plausibility_review"),
|
||||
_StubReviewStage("antonym_reversal_review"),
|
||||
ProtectedGlossaryTermsValidator("protected_glossary_guard"),
|
||||
StubLLMValidator("acoustic_similarity_review"),
|
||||
StubLLMValidator("contextual_plausibility_review"),
|
||||
StubLLMValidator("antonym_reversal_review"),
|
||||
]
|
||||
|
||||
def propose(
|
||||
|
||||
@@ -1,61 +1,19 @@
|
||||
from typing import Sequence
|
||||
|
||||
from audita.core.config import AuditaConfig
|
||||
from audita.core.schemas import Glossary, TranscriptSegment
|
||||
from audita.framework.models import (
|
||||
CorrectionProposal,
|
||||
DeterministicFilter,
|
||||
FilterDecision,
|
||||
ModuleContext,
|
||||
ReviewDecision,
|
||||
ReviewStage,
|
||||
)
|
||||
|
||||
|
||||
class _StubFilter:
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
proposal: CorrectionProposal,
|
||||
transcript: Sequence[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
) -> FilterDecision:
|
||||
return FilterDecision(approved=True)
|
||||
|
||||
|
||||
class _StubReviewStage:
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
|
||||
def review(
|
||||
self,
|
||||
proposals: Sequence[CorrectionProposal],
|
||||
transcript: Sequence[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
llm_client,
|
||||
run_dir,
|
||||
) -> Sequence[ReviewDecision]:
|
||||
return [ReviewDecision(proposal_index=proposal.proposal_index, approved=True) for proposal in proposals]
|
||||
from audita.core.schemas import TranscriptSegment
|
||||
from audita.framework.models import CorrectionProposal, ModuleContext
|
||||
from audita.validators import ProtectedGlossaryTermsValidator, StubLLMValidator, Validator
|
||||
|
||||
|
||||
class SpokenWordModule:
|
||||
module_key = "spoken_word"
|
||||
replacement_policy = "replace_all"
|
||||
|
||||
def deterministic_filters(self) -> Sequence[DeterministicFilter]:
|
||||
def validators(self) -> Sequence[Validator]:
|
||||
return [
|
||||
_StubFilter("protected_glossary_guard"),
|
||||
_StubFilter("spoken_disfluency_guard"),
|
||||
]
|
||||
|
||||
def review_stages(self) -> Sequence[ReviewStage]:
|
||||
return [
|
||||
_StubReviewStage("spoken_marker_cleanup_review"),
|
||||
_StubReviewStage("meaning_preservation_review"),
|
||||
ProtectedGlossaryTermsValidator("protected_glossary_guard"),
|
||||
StubLLMValidator("spoken_marker_cleanup_review"),
|
||||
StubLLMValidator("meaning_preservation_review"),
|
||||
]
|
||||
|
||||
def propose(
|
||||
|
||||
14
src/audita/validators/__init__.py
Normal file
14
src/audita/validators/__init__.py
Normal file
@@ -0,0 +1,14 @@
|
||||
from .base import ValidationContext, ValidationDecision, ValidationResult, Validator
|
||||
from .deterministic import ProtectedGlossaryTermsValidator
|
||||
from .llm import StubLLMValidator
|
||||
from .protection import ProtectedVocabulary
|
||||
|
||||
__all__ = [
|
||||
"ValidationContext",
|
||||
"ValidationDecision",
|
||||
"ValidationResult",
|
||||
"Validator",
|
||||
"ProtectedGlossaryTermsValidator",
|
||||
"ProtectedVocabulary",
|
||||
"StubLLMValidator",
|
||||
]
|
||||
45
src/audita/validators/base.py
Normal file
45
src/audita/validators/base.py
Normal file
@@ -0,0 +1,45 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, List, Optional, Protocol, Sequence
|
||||
|
||||
from audita.core.config import AuditaConfig
|
||||
from audita.core.schemas import Glossary, TranscriptSegment
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from audita.framework.models import CorrectionProposal, ModuleRunSpec, StructuredLLMClient
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ValidationContext:
|
||||
proposals: Sequence["CorrectionProposal"]
|
||||
transcript: Sequence[TranscriptSegment]
|
||||
glossary: Glossary
|
||||
config: AuditaConfig
|
||||
run_spec: "ModuleRunSpec"
|
||||
run_dir: Path
|
||||
llm_client: Optional["StructuredLLMClient"] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ValidationDecision:
|
||||
proposal_index: int
|
||||
approved: bool
|
||||
confidence: Optional[float] = None
|
||||
reason: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ValidationResult:
|
||||
validator_name: str
|
||||
execution_kind: str
|
||||
decisions: List[ValidationDecision]
|
||||
|
||||
|
||||
class Validator(Protocol):
|
||||
name: str
|
||||
execution_kind: str
|
||||
|
||||
def validate(self, context: ValidationContext) -> ValidationResult:
|
||||
...
|
||||
28
src/audita/validators/deterministic.py
Normal file
28
src/audita/validators/deterministic.py
Normal file
@@ -0,0 +1,28 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .base import ValidationContext, ValidationDecision, ValidationResult
|
||||
from .protection import ProtectedVocabulary
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProtectedGlossaryTermsValidator:
|
||||
name: str
|
||||
execution_kind: str = "deterministic"
|
||||
|
||||
def validate(self, context: ValidationContext) -> ValidationResult:
|
||||
vocabulary = ProtectedVocabulary.from_glossary(context.glossary)
|
||||
return ValidationResult(
|
||||
validator_name=self.name,
|
||||
execution_kind=self.execution_kind,
|
||||
decisions=[
|
||||
ValidationDecision(
|
||||
proposal_index=proposal.proposal_index,
|
||||
approved=(reason := vocabulary.violation_reason(
|
||||
proposal.original_text,
|
||||
proposal.corrected_text,
|
||||
)) is None,
|
||||
reason=reason,
|
||||
)
|
||||
for proposal in context.proposals
|
||||
],
|
||||
)
|
||||
19
src/audita/validators/llm.py
Normal file
19
src/audita/validators/llm.py
Normal file
@@ -0,0 +1,19 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .base import ValidationContext, ValidationDecision, ValidationResult
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StubLLMValidator:
|
||||
name: str
|
||||
execution_kind: str = "llm"
|
||||
|
||||
def validate(self, context: ValidationContext) -> ValidationResult:
|
||||
return ValidationResult(
|
||||
validator_name=self.name,
|
||||
execution_kind=self.execution_kind,
|
||||
decisions=[
|
||||
ValidationDecision(proposal_index=proposal.proposal_index, approved=True)
|
||||
for proposal in context.proposals
|
||||
],
|
||||
)
|
||||
118
src/audita/validators/protection.py
Normal file
118
src/audita/validators/protection.py
Normal file
@@ -0,0 +1,118 @@
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional, Pattern
|
||||
|
||||
from audita.core.schemas import Glossary
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProtectedVocabulary:
|
||||
terms_by_folded: Dict[str, "_ProtectedTermDefinition"]
|
||||
pattern: Optional[Pattern[str]]
|
||||
|
||||
@classmethod
|
||||
def from_glossary(cls, glossary: Glossary) -> "ProtectedVocabulary":
|
||||
terms_by_folded: Dict[str, _ProtectedTermDefinition] = {}
|
||||
for identity, entry in enumerate(glossary.glossary):
|
||||
entry_terms = [entry.name, *entry.aliases]
|
||||
for term in entry_terms:
|
||||
_add_term(terms_by_folded, term, identity)
|
||||
_add_term(terms_by_folded, f"{term}s", identity)
|
||||
if entry.plural is not None:
|
||||
_add_term(terms_by_folded, entry.plural, identity)
|
||||
|
||||
terms = [definition.canonical for definition in terms_by_folded.values()]
|
||||
if not terms:
|
||||
return cls(terms_by_folded=terms_by_folded, pattern=None)
|
||||
|
||||
alternatives = sorted((re.escape(term) for term in terms), key=len, reverse=True)
|
||||
pattern = re.compile(r"(?<!\w)(" + "|".join(alternatives) + r")(?!\w)", flags=re.IGNORECASE)
|
||||
return cls(terms_by_folded=terms_by_folded, pattern=pattern)
|
||||
|
||||
def violation_reason(self, before: str, after: str) -> Optional[str]:
|
||||
before_occurrences = self._occurrences_by_identity(before)
|
||||
after_occurrences = self._occurrences_by_identity(after)
|
||||
|
||||
reason = self._validate_identity_preservation(before_occurrences, after_occurrences)
|
||||
if reason is not None:
|
||||
return reason
|
||||
return self._validate_capitalization_transitions(before_occurrences, after_occurrences)
|
||||
|
||||
def _occurrences(self, text: str) -> List["_ProtectedOccurrence"]:
|
||||
if self.pattern is None:
|
||||
return []
|
||||
occurrences = []
|
||||
for match in self.pattern.finditer(text):
|
||||
matched_text = match.group(0)
|
||||
definition = self.terms_by_folded[matched_text.casefold()]
|
||||
occurrences.append(
|
||||
_ProtectedOccurrence(
|
||||
text=matched_text,
|
||||
identity=definition.identity,
|
||||
canonical=definition.canonical,
|
||||
)
|
||||
)
|
||||
return occurrences
|
||||
|
||||
def _occurrences_by_identity(self, text: str) -> Dict[int, List["_ProtectedOccurrence"]]:
|
||||
occurrences_by_identity: Dict[int, List["_ProtectedOccurrence"]] = {}
|
||||
for occurrence in self._occurrences(text):
|
||||
occurrences_by_identity.setdefault(occurrence.identity, []).append(occurrence)
|
||||
return occurrences_by_identity
|
||||
|
||||
def _validate_identity_preservation(
|
||||
self,
|
||||
before_occurrences: Dict[int, List["_ProtectedOccurrence"]],
|
||||
after_occurrences: Dict[int, List["_ProtectedOccurrence"]],
|
||||
) -> Optional[str]:
|
||||
for identity, before_items in before_occurrences.items():
|
||||
if len(after_occurrences.get(identity, [])) < len(before_items):
|
||||
return "correction changes protected glossary term usage"
|
||||
return None
|
||||
|
||||
def _validate_capitalization_transitions(
|
||||
self,
|
||||
before_occurrences: Dict[int, List["_ProtectedOccurrence"]],
|
||||
after_occurrences: Dict[int, List["_ProtectedOccurrence"]],
|
||||
) -> Optional[str]:
|
||||
for identity, after_items in after_occurrences.items():
|
||||
before_items = before_occurrences.get(identity, [])
|
||||
before_count = len(before_items)
|
||||
for index, after_item in enumerate(after_items):
|
||||
if index < before_count:
|
||||
before_item = before_items[index]
|
||||
if after_item.text == before_item.text:
|
||||
continue
|
||||
if after_item.text == after_item.canonical:
|
||||
continue
|
||||
return "correction changes protected glossary term capitalization"
|
||||
if after_item.text != after_item.canonical:
|
||||
return "correction changes protected glossary term capitalization"
|
||||
return None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ProtectedOccurrence:
|
||||
text: str
|
||||
identity: int
|
||||
canonical: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ProtectedTermDefinition:
|
||||
identity: int
|
||||
canonical: str
|
||||
|
||||
|
||||
def _add_term(
|
||||
terms_by_folded: Dict[str, _ProtectedTermDefinition],
|
||||
term: str,
|
||||
identity: int,
|
||||
) -> None:
|
||||
stripped = term.strip()
|
||||
if not stripped:
|
||||
return
|
||||
terms_by_folded.setdefault(
|
||||
stripped.casefold(),
|
||||
_ProtectedTermDefinition(identity=identity, canonical=stripped),
|
||||
)
|
||||
@@ -1,81 +1,56 @@
|
||||
from pathlib import Path
|
||||
|
||||
from audita.core.config import AuditaConfig
|
||||
from audita.core.schemas import parse_glossary_yaml, parse_transcript_json
|
||||
from audita.framework.models import (
|
||||
CorrectionProposal,
|
||||
FilterDecision,
|
||||
ModuleContext,
|
||||
ModuleRunSpec,
|
||||
ReviewDecision,
|
||||
)
|
||||
from audita.framework.models import CorrectionProposal, ModuleContext, ModuleRunSpec
|
||||
from audita.framework.runner import PipelineRunner
|
||||
from audita.validators import ProtectedGlossaryTermsValidator
|
||||
from audita.validators.base import ValidationContext, ValidationDecision, ValidationResult
|
||||
|
||||
|
||||
class AllowAllFilter:
|
||||
name = "allow_all"
|
||||
class RecordingValidator:
|
||||
execution_kind = "deterministic"
|
||||
|
||||
def evaluate(self, proposal, transcript, glossary, config):
|
||||
return FilterDecision(approved=True)
|
||||
def __init__(self, name, recorder, approve=True):
|
||||
self.name = name
|
||||
self._recorder = recorder
|
||||
self._approve = approve
|
||||
|
||||
def validate(self, context: ValidationContext) -> ValidationResult:
|
||||
self._recorder.append((self.name, [proposal.corrected_text for proposal in context.proposals]))
|
||||
return ValidationResult(
|
||||
validator_name=self.name,
|
||||
execution_kind=self.execution_kind,
|
||||
decisions=[
|
||||
ValidationDecision(
|
||||
proposal_index=proposal.proposal_index,
|
||||
approved=self._approve,
|
||||
reason=None if self._approve else f"{self.name} rejected proposal",
|
||||
)
|
||||
for proposal in context.proposals
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
class RejectAllFilter:
|
||||
name = "reject_all"
|
||||
|
||||
def evaluate(self, proposal, transcript, glossary, config):
|
||||
return FilterDecision(approved=False, reason="filter rejected proposal")
|
||||
|
||||
|
||||
class AllowAllReviewStage:
|
||||
name = "allow_all_review"
|
||||
|
||||
def review(self, proposals, transcript, glossary, config, llm_client, run_dir):
|
||||
return [ReviewDecision(proposal_index=proposal.proposal_index, approved=True) for proposal in proposals]
|
||||
class RecordingLLMValidator(RecordingValidator):
|
||||
execution_kind = "llm"
|
||||
|
||||
|
||||
class RecordingModule:
|
||||
replacement_policy = "require_unique"
|
||||
|
||||
def __init__(self, module_key, proposals, recorder):
|
||||
def __init__(self, module_key, proposals, validators, recorder):
|
||||
self.module_key = module_key
|
||||
self._proposals = proposals
|
||||
self._validators = validators
|
||||
self._recorder = recorder
|
||||
|
||||
def deterministic_filters(self):
|
||||
return [AllowAllFilter()]
|
||||
|
||||
def review_stages(self):
|
||||
return [AllowAllReviewStage()]
|
||||
def validators(self):
|
||||
return list(self._validators)
|
||||
|
||||
def propose(self, transcript_section, context: ModuleContext):
|
||||
self._recorder.append([segment.text for segment in transcript_section])
|
||||
self._recorder.append(("propose", [segment.text for segment in transcript_section]))
|
||||
return list(self._proposals)
|
||||
|
||||
|
||||
class RejectedModule:
|
||||
module_key = "rejected"
|
||||
replacement_policy = "require_unique"
|
||||
|
||||
def deterministic_filters(self):
|
||||
return [RejectAllFilter()]
|
||||
|
||||
def review_stages(self):
|
||||
return []
|
||||
|
||||
def propose(self, transcript_section, context):
|
||||
return [
|
||||
CorrectionProposal(
|
||||
proposal_index=0,
|
||||
module_instance=context.run_spec.instance_name,
|
||||
module_key=context.run_spec.module_key,
|
||||
id=1,
|
||||
original_text="Hello",
|
||||
corrected_text="Goodbye",
|
||||
confidence=0.9,
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def test_pipeline_runner_applies_modules_sequentially(tmp_path):
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
@@ -106,6 +81,7 @@ def test_pipeline_runner_applies_modules_sequentially(tmp_path):
|
||||
confidence=0.9,
|
||||
)
|
||||
],
|
||||
[RecordingValidator("first_validator", seen)],
|
||||
seen,
|
||||
)
|
||||
second = RecordingModule(
|
||||
@@ -121,6 +97,7 @@ def test_pipeline_runner_applies_modules_sequentially(tmp_path):
|
||||
confidence=0.9,
|
||||
)
|
||||
],
|
||||
[RecordingValidator("second_validator", seen)],
|
||||
seen,
|
||||
)
|
||||
|
||||
@@ -136,13 +113,15 @@ def test_pipeline_runner_applies_modules_sequentially(tmp_path):
|
||||
run_dir=tmp_path / "run",
|
||||
)
|
||||
|
||||
assert seen[0] == ["Alpha."]
|
||||
assert seen[1] == ["Beta."]
|
||||
assert seen[0] == ("propose", ["Alpha."])
|
||||
assert seen[1] == ("first_validator", ["Beta"])
|
||||
assert seen[2] == ("propose", ["Beta."])
|
||||
assert seen[3] == ("second_validator", ["Gamma"])
|
||||
assert result.transcript[0].text == "Gamma."
|
||||
assert len(result.applied_changes) == 2
|
||||
|
||||
|
||||
def test_pipeline_runner_reports_filter_rejections(tmp_path):
|
||||
def test_pipeline_runner_validator_order_respects_survivors(tmp_path):
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
@@ -158,16 +137,138 @@ def test_pipeline_runner_reports_filter_rejections(tmp_path):
|
||||
summary: "Hello."
|
||||
"""
|
||||
)
|
||||
seen = []
|
||||
module = RecordingModule(
|
||||
"mod",
|
||||
[
|
||||
CorrectionProposal(
|
||||
proposal_index=0,
|
||||
module_instance="mod",
|
||||
module_key="mod",
|
||||
id=1,
|
||||
original_text="Hello",
|
||||
corrected_text="Goodbye",
|
||||
confidence=0.9,
|
||||
)
|
||||
],
|
||||
[
|
||||
RecordingValidator("first", seen, approve=False),
|
||||
RecordingLLMValidator("second", seen, approve=True),
|
||||
],
|
||||
seen,
|
||||
)
|
||||
|
||||
runner = PipelineRunner()
|
||||
result = runner.run(
|
||||
transcript=transcript,
|
||||
glossary=glossary,
|
||||
module_specs=[ModuleRunSpec(instance_name="rejected", module_key="rejected", module=RejectedModule())],
|
||||
module_specs=[ModuleRunSpec(instance_name="mod", module_key="mod", module=module)],
|
||||
config=AuditaConfig.from_sources(env={}),
|
||||
run_dir=tmp_path / "run",
|
||||
)
|
||||
|
||||
assert result.transcript[0].text == "Hello."
|
||||
assert result.module_reports[0].skipped_count == 1
|
||||
assert result.skipped_corrections[0].reason == "filter rejected proposal"
|
||||
assert ("first", ["Goodbye"]) in seen
|
||||
assert all(entry[0] != "second" for entry in seen)
|
||||
assert result.module_reports[0].validators[0].rejected_count == 1
|
||||
assert result.module_reports[0].validators[1].candidate_count == 0
|
||||
assert result.skipped_corrections[0].source == "validator:first"
|
||||
|
||||
|
||||
def test_pipeline_runner_supports_deterministic_and_llm_validators_in_one_chain(tmp_path):
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "A", "start": 0.0, "end": 1.0, "text": "Alpha."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Alpha"
|
||||
category: noun
|
||||
summary: "Alpha."
|
||||
"""
|
||||
)
|
||||
seen = []
|
||||
module = RecordingModule(
|
||||
"mixed",
|
||||
[
|
||||
CorrectionProposal(
|
||||
proposal_index=0,
|
||||
module_instance="mixed",
|
||||
module_key="mixed",
|
||||
id=1,
|
||||
original_text="Alpha",
|
||||
corrected_text="Beta",
|
||||
confidence=0.9,
|
||||
)
|
||||
],
|
||||
[
|
||||
RecordingValidator("deterministic_guard", seen),
|
||||
RecordingLLMValidator("llm_review", seen),
|
||||
],
|
||||
seen,
|
||||
)
|
||||
|
||||
runner = PipelineRunner()
|
||||
result = runner.run(
|
||||
transcript=transcript,
|
||||
glossary=glossary,
|
||||
module_specs=[ModuleRunSpec(instance_name="mixed", module_key="mixed", module=module)],
|
||||
config=AuditaConfig.from_sources(env={}),
|
||||
run_dir=tmp_path / "run",
|
||||
)
|
||||
|
||||
assert result.transcript[0].text == "Beta."
|
||||
assert [report.execution_kind for report in result.module_reports[0].validators] == [
|
||||
"deterministic",
|
||||
"llm",
|
||||
]
|
||||
|
||||
|
||||
def test_pipeline_runner_uses_real_protected_glossary_validator(tmp_path):
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "A", "start": 0.0, "end": 1.0, "text": "Hrank moves."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Hrank"
|
||||
category: pc
|
||||
summary: "Hrank is a player character."
|
||||
"""
|
||||
)
|
||||
module = RecordingModule(
|
||||
"protected",
|
||||
[
|
||||
CorrectionProposal(
|
||||
proposal_index=0,
|
||||
module_instance="protected",
|
||||
module_key="protected",
|
||||
id=1,
|
||||
original_text="Hrank",
|
||||
corrected_text="Frank",
|
||||
confidence=0.9,
|
||||
)
|
||||
],
|
||||
[ProtectedGlossaryTermsValidator("protected_glossary_guard")],
|
||||
[],
|
||||
)
|
||||
|
||||
runner = PipelineRunner()
|
||||
result = runner.run(
|
||||
transcript=transcript,
|
||||
glossary=glossary,
|
||||
module_specs=[ModuleRunSpec(instance_name="protected", module_key="protected", module=module)],
|
||||
config=AuditaConfig.from_sources(env={}),
|
||||
run_dir=tmp_path / "run",
|
||||
)
|
||||
|
||||
assert result.transcript[0].text == "Hrank moves."
|
||||
assert result.skipped_corrections[0].source == "validator:protected_glossary_guard"
|
||||
assert result.skipped_corrections[0].reason == "correction changes protected glossary term usage"
|
||||
|
||||
@@ -3,6 +3,7 @@ import json
|
||||
from audita.core.config import AuditaConfig
|
||||
from audita.core.io import write_report
|
||||
from audita.core.schemas import parse_glossary_yaml, parse_source_transcript_json
|
||||
from audita.modules import default_module_specs
|
||||
from audita.pipeline import process_transcript, process_transcript_result
|
||||
|
||||
|
||||
@@ -73,6 +74,11 @@ def test_process_transcript_result_writes_report_and_preserves_skips_per_policy(
|
||||
assert result.report.totals["applied_change_count"] == 0
|
||||
assert (result.run_dir / "report.json").exists()
|
||||
assert (result.run_dir / "normalization" / "summary.json").exists()
|
||||
assert [validator["name"] for validator in result.report.modules[0].to_dict()["validators"]] == [
|
||||
"protected_glossary_guard",
|
||||
"toward_glossary_term_review",
|
||||
"context_support_review",
|
||||
]
|
||||
|
||||
|
||||
def test_external_report_can_be_written(tmp_path):
|
||||
@@ -84,3 +90,33 @@ def test_external_report_can_be_written(tmp_path):
|
||||
payload = json.loads(report_path.read_text(encoding="utf-8"))
|
||||
assert payload["pipeline"][0] == "glossary_primary"
|
||||
assert payload["totals"]["applied_change_count"] == 0
|
||||
|
||||
|
||||
def test_default_module_specs_expose_final_validator_order():
|
||||
specs = default_module_specs()
|
||||
|
||||
assert [validator.name for validator in specs[0].module.validators()] == [
|
||||
"protected_glossary_guard",
|
||||
"toward_glossary_term_review",
|
||||
"context_support_review",
|
||||
]
|
||||
assert [validator.name for validator in specs[1].module.validators()] == [
|
||||
"protected_glossary_guard",
|
||||
"acoustic_similarity_review",
|
||||
"contextual_plausibility_review",
|
||||
"antonym_reversal_review",
|
||||
]
|
||||
assert [validator.name for validator in specs[2].module.validators()] == [
|
||||
"protected_glossary_guard",
|
||||
"toward_glossary_term_review",
|
||||
"context_support_review",
|
||||
]
|
||||
assert [validator.name for validator in specs[3].module.validators()] == [
|
||||
"protected_glossary_guard",
|
||||
"spoken_marker_cleanup_review",
|
||||
"meaning_preservation_review",
|
||||
]
|
||||
assert [validator.name for validator in specs[4].module.validators()] == [
|
||||
"protected_glossary_guard",
|
||||
"edited_text_readability_review",
|
||||
]
|
||||
|
||||
201
tests/test_protected_validator.py
Normal file
201
tests/test_protected_validator.py
Normal file
@@ -0,0 +1,201 @@
|
||||
from audita.core.schemas import parse_glossary_yaml, parse_transcript_json
|
||||
from audita.framework.models import CorrectionProposal, ModuleRunSpec
|
||||
from audita.validators import ProtectedGlossaryTermsValidator, ProtectedVocabulary
|
||||
from audita.validators.base import ValidationContext
|
||||
|
||||
|
||||
def _glossary():
|
||||
return parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Hrank"
|
||||
aliases:
|
||||
- "Greenfield"
|
||||
category: pc
|
||||
summary: "Hrank Greenfield is a player character."
|
||||
- name: "Popov"
|
||||
category: npc
|
||||
summary: "Popov is an allied NPC."
|
||||
- name: "Jesters"
|
||||
aliases:
|
||||
- "Jester"
|
||||
category: faction
|
||||
summary: "The Jesters are a faction."
|
||||
- name: "Svend"
|
||||
category: pc
|
||||
summary: "Svend is a player character."
|
||||
- name: "Godfrey"
|
||||
category: npc
|
||||
summary: "Godfrey is an NPC."
|
||||
- name: "Lyra"
|
||||
category: npc
|
||||
summary: "Lyra is an NPC."
|
||||
- name: "Loviator"
|
||||
category: deity
|
||||
summary: "Loviator is a deity."
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_protected_vocabulary_blocks_replacing_protected_term():
|
||||
vocabulary = ProtectedVocabulary.from_glossary(_glossary())
|
||||
|
||||
assert (
|
||||
vocabulary.violation_reason("Hrank moves.", "Frank moves.")
|
||||
== "correction changes protected glossary term usage"
|
||||
)
|
||||
|
||||
|
||||
def test_protected_vocabulary_blocks_noncanonical_capitalization():
|
||||
vocabulary = ProtectedVocabulary.from_glossary(_glossary())
|
||||
|
||||
assert (
|
||||
vocabulary.violation_reason("Popov moves.", "POPOV moves.")
|
||||
== "correction changes protected glossary term capitalization"
|
||||
)
|
||||
assert (
|
||||
vocabulary.violation_reason("gestures", "jesters")
|
||||
== "correction changes protected glossary term capitalization"
|
||||
)
|
||||
|
||||
|
||||
def test_protected_vocabulary_allows_corrections_toward_protected_terms():
|
||||
vocabulary = ProtectedVocabulary.from_glossary(_glossary())
|
||||
|
||||
assert vocabulary.violation_reason("Pawpaw moves.", "Popov moves.") is None
|
||||
assert vocabulary.violation_reason("gestures", "Jesters") is None
|
||||
assert vocabulary.violation_reason("rank", "Hrank") is None
|
||||
assert vocabulary.violation_reason("spend", "Svend") is None
|
||||
|
||||
|
||||
def test_protected_vocabulary_allows_unchanged_noncanonical_terms_and_quote_wrapping():
|
||||
vocabulary = ProtectedVocabulary.from_glossary(_glossary())
|
||||
|
||||
assert vocabulary.violation_reason("jesters advance.", "jesters advance.") is None
|
||||
before = (
|
||||
"When you say that, Popov will say, when I was in that room with the jesters, "
|
||||
"I just knew that Godfrey and Lyra came directly from Loviator herself."
|
||||
)
|
||||
after = (
|
||||
'When you say that, Popov will say, "When I was in that room with the jesters, '
|
||||
'I just knew that Godfrey and Lyra came directly from Loviator herself."'
|
||||
)
|
||||
assert vocabulary.violation_reason(before, after) is None
|
||||
|
||||
|
||||
def test_protected_vocabulary_allows_inferred_and_explicit_plurals():
|
||||
vocabulary = ProtectedVocabulary.from_glossary(_glossary())
|
||||
explicit = ProtectedVocabulary.from_glossary(
|
||||
parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Mox"
|
||||
plural: "Moxen"
|
||||
category: faction
|
||||
summary: "The Mox are a faction."
|
||||
"""
|
||||
)
|
||||
)
|
||||
|
||||
assert vocabulary.violation_reason("Godfrey's", "Godfreys") is None
|
||||
assert vocabulary.violation_reason("gesture", "Jesters") is None
|
||||
assert explicit.violation_reason("Mox's", "Moxen") is None
|
||||
|
||||
|
||||
def test_protected_vocabulary_does_not_match_embedded_substrings():
|
||||
vocabulary = ProtectedVocabulary.from_glossary(_glossary())
|
||||
|
||||
assert vocabulary.violation_reason("The shrank spell worked.", "The shrank spell works.") is None
|
||||
|
||||
|
||||
def test_protected_glossary_terms_validator_returns_proposal_indexed_decisions():
|
||||
validator = ProtectedGlossaryTermsValidator("protected_glossary_guard")
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Frank moves."},
|
||||
{"id": 2, "speaker": "Eric", "start": 1.0, "end": 2.0, "text": "Pawpaw waits."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
proposals = [
|
||||
CorrectionProposal(
|
||||
proposal_index=0,
|
||||
module_instance="glossary_primary",
|
||||
module_key="glossary",
|
||||
id=1,
|
||||
original_text="Hrank",
|
||||
corrected_text="Frank",
|
||||
confidence=0.9,
|
||||
),
|
||||
CorrectionProposal(
|
||||
proposal_index=1,
|
||||
module_instance="glossary_primary",
|
||||
module_key="glossary",
|
||||
id=2,
|
||||
original_text="Pawpaw",
|
||||
corrected_text="Popov",
|
||||
confidence=0.9,
|
||||
),
|
||||
]
|
||||
result = validator.validate(
|
||||
ValidationContext(
|
||||
proposals=proposals,
|
||||
transcript=transcript,
|
||||
glossary=_glossary(),
|
||||
config=None, # type: ignore[arg-type]
|
||||
run_spec=ModuleRunSpec(instance_name="glossary_primary", module_key="glossary", module=None), # type: ignore[arg-type]
|
||||
run_dir=transcript[0].__class__.__module__ and __import__("pathlib").Path("."),
|
||||
)
|
||||
)
|
||||
|
||||
assert [decision.proposal_index for decision in result.decisions] == [0, 1]
|
||||
assert result.decisions[0].approved is False
|
||||
assert result.decisions[0].reason == "correction changes protected glossary term usage"
|
||||
assert result.decisions[1].approved is True
|
||||
|
||||
|
||||
def test_protected_glossary_terms_validator_uses_proposal_span_only():
|
||||
validator = ProtectedGlossaryTermsValidator("protected_glossary_guard")
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Svend"
|
||||
category: pc
|
||||
summary: "Svend is a player character."
|
||||
- name: "Jesters"
|
||||
category: faction
|
||||
summary: "The Jesters are a faction."
|
||||
"""
|
||||
)
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "You have to keep it bind. Svend sees the jesters."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
proposals = [
|
||||
CorrectionProposal(
|
||||
proposal_index=0,
|
||||
module_instance="spoken_word",
|
||||
module_key="spoken_word",
|
||||
id=1,
|
||||
original_text="keep it bind",
|
||||
corrected_text="keep in mind",
|
||||
confidence=0.9,
|
||||
)
|
||||
]
|
||||
|
||||
result = validator.validate(
|
||||
ValidationContext(
|
||||
proposals=proposals,
|
||||
transcript=transcript,
|
||||
glossary=glossary,
|
||||
config=None, # type: ignore[arg-type]
|
||||
run_spec=ModuleRunSpec(instance_name="spoken_word", module_key="spoken_word", module=None), # type: ignore[arg-type]
|
||||
run_dir=__import__("pathlib").Path("."),
|
||||
)
|
||||
)
|
||||
|
||||
assert result.decisions[0].approved is True
|
||||
Reference in New Issue
Block a user