Compare commits
62 Commits
main
..
docs/polish
| Author | SHA1 | Date | |
|---|---|---|---|
|
0699e97e1d
|
|||
|
7de4505e31
|
|||
|
564d663c78
|
|||
|
0b2bc6c688
|
|||
|
aa687f43cd
|
|||
|
f18427e123
|
|||
|
1754556c99
|
|||
|
13c869f5d6
|
|||
|
93515707cc
|
|||
|
3cde5aba77
|
|||
|
69966d171c
|
|||
|
d5df8b52e6
|
|||
|
8b7c087de7
|
|||
|
c54f8c3f6f
|
|||
|
0cadfd4d23
|
|||
|
e128720917
|
|||
|
f713d5e8a6
|
|||
|
9853b57843
|
|||
|
55faae3e1b
|
|||
|
07ac70e835
|
|||
|
6d1bece815
|
|||
|
40fa39485e
|
|||
|
9115f0c25b
|
|||
|
83c4b4bee5
|
|||
|
44e3e8f4ea
|
|||
|
45dfe07772
|
|||
|
6bafc43754
|
|||
|
012b306749
|
|||
|
ac7c5c69cf
|
|||
|
cd36c54e22
|
|||
|
107fc4e275
|
|||
|
571b770281
|
|||
|
8b3536873e
|
|||
|
9a4ac359a3
|
|||
|
de5ad93524
|
|||
|
cab8099d06
|
|||
|
e2be3daffd
|
|||
|
9239300fd9
|
|||
|
b9f805b2f4
|
|||
|
75cd7b68de
|
|||
|
b2b5eb1518
|
|||
|
9e7b0131b3
|
|||
|
b8ab5811dd
|
|||
|
62fac688e4
|
|||
|
14ac7dbbb9
|
|||
|
aad933f9c4
|
|||
|
2a7476046d
|
|||
|
914c738013
|
|||
|
a4f5fa4cc6
|
|||
|
027d2d3beb
|
|||
|
74ee8c2e7b
|
|||
|
e1c8c25142
|
|||
|
e6167005e5
|
|||
|
14dcddcbba
|
|||
|
1e3618e637
|
|||
|
a65249fa44
|
|||
|
697b1ddfeb
|
|||
|
efc6a031a3
|
|||
|
a1e862550c
|
|||
|
60aaa33327
|
|||
|
818e241ab2
|
|||
|
49f1e27cb1
|
@@ -0,0 +1,121 @@
|
||||
# CLAUDE.md
|
||||
|
||||
Guidelines for working on the Veritext project.
|
||||
|
||||
## Project Overview
|
||||
|
||||
Veritext is a semantic text validation framework for Python. It validates text outputs
|
||||
against quality criteria using metrics like BLEU, ROUGE, and semantic similarity.
|
||||
|
||||
## Directory Structure
|
||||
|
||||
```
|
||||
veritext/
|
||||
├── src/veritext/ # Package source
|
||||
│ ├── core/ # Shared types, tokenisation, config
|
||||
│ ├── metrics/ # BLEU, ROUGE, lexical, readability
|
||||
│ ├── semantic/ # Optional embedding-based similarity
|
||||
│ ├── validators/ # Composable validation checks
|
||||
│ ├── benchmark/ # Quality tracking & regression detection
|
||||
│ ├── pytest_plugin/ # Native pytest integration
|
||||
│ └── cli/ # Command-line interface
|
||||
├── tests/ # Test suite (mirrors src structure)
|
||||
├── docs/ # Project documentation
|
||||
└── examples/ # Usage examples
|
||||
```
|
||||
|
||||
## Code Style
|
||||
|
||||
### Python Conventions
|
||||
|
||||
- **Python 3.11+** with modern type hints
|
||||
- **UK English** in all text (colour, behaviour, summarisation, tokenisation)
|
||||
- **snake_case** for variables, functions, modules
|
||||
- **PascalCase** for classes
|
||||
- Absolute imports from package root: `from veritext.core.types import ...`
|
||||
|
||||
### Quality Gates
|
||||
|
||||
All must pass with zero issues before any commit:
|
||||
|
||||
```bash
|
||||
uv run ruff check . # Linting
|
||||
uv run ruff format --check . # Formatting
|
||||
uv run mypy src/ # Type checking
|
||||
uv run pytest # Tests
|
||||
```
|
||||
|
||||
### Documentation
|
||||
|
||||
- Docstrings for all public APIs (Google style)
|
||||
- Type hints on all function signatures
|
||||
- Keep docstrings concise; let types speak where possible
|
||||
|
||||
## Architecture
|
||||
|
||||
### Layer Dependencies
|
||||
|
||||
```
|
||||
CLI / pytest_plugin (presentation)
|
||||
↓
|
||||
validators / benchmark (decision logic)
|
||||
↓
|
||||
metrics (pure computation)
|
||||
↓
|
||||
core (shared types, tokenisation)
|
||||
```
|
||||
|
||||
Each layer depends only on layers below it.
|
||||
|
||||
### Metrics vs Validators
|
||||
|
||||
| Concept | Responsibility | Output |
|
||||
|---------|----------------|--------|
|
||||
| **Metric** | Compute a score | Typed result (e.g., `BleuResult`) |
|
||||
| **Validator** | Make pass/fail decision | `ValidationResult` with diagnostics |
|
||||
|
||||
### Edge Case Handling
|
||||
|
||||
- Empty text: Metrics return zero scores; validators fail
|
||||
- Empty reference: Comparison metrics raise `ValueError`
|
||||
- Whitespace-only: Treated as empty after tokenisation
|
||||
- Unicode: NFC normalisation by default
|
||||
|
||||
## Git Workflow
|
||||
|
||||
### Before Starting Work
|
||||
|
||||
When starting work from a plan, create a new branch matching the plan's scope before
|
||||
making any changes. Do not reuse an existing branch from previous work, even if related.
|
||||
|
||||
### Commits
|
||||
|
||||
- Format: `type(scope): description`
|
||||
- Types: feat, fix, chore, refactor, docs, test
|
||||
- Atomic: ≤3 new files, ≤150 LOC per commit
|
||||
- Update changelog.md before completing a task
|
||||
|
||||
### Branches
|
||||
|
||||
- `feat/kebab-case` — new features
|
||||
- `fix/kebab-case` — bug fixes
|
||||
- `chore/` — maintenance
|
||||
- `refactor/` — code restructure
|
||||
- `docs/` — documentation only
|
||||
|
||||
## Testing
|
||||
|
||||
- Test files mirror source structure: `tests/test_core/test_types.py`
|
||||
- Use pytest fixtures for common setup
|
||||
- Target ≥80% coverage
|
||||
- Include edge cases: empty input, Unicode, boundary values
|
||||
|
||||
## Pre-Completion Checklist
|
||||
|
||||
Before marking ANY task complete:
|
||||
|
||||
- [ ] All linting/formatting/type checks pass
|
||||
- [ ] Tests pass with adequate coverage
|
||||
- [ ] changelog.md updated if user-facing changes
|
||||
- [ ] Filenames are lowercase (except CLAUDE.md)
|
||||
- [ ] Commit follows `type(scope): description` format
|
||||
+2
-2
@@ -36,7 +36,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
- Added `cache_max_size` parameter to `SemanticSimilarity` (default: 1000 embeddings)
|
||||
- Added test coverage for `core/config.py` and `core/logging.py` modules
|
||||
|
||||
## [0.1.0] — 2025-05-17
|
||||
## [0.1.0] — 2026-02-03
|
||||
|
||||
Initial release of Veritext, a semantic text validation framework for Python.
|
||||
|
||||
@@ -110,5 +110,5 @@ Initial release of Veritext, a semantic text validation framework for Python.
|
||||
|
||||
#### Documentation
|
||||
|
||||
- Readme with usage examples
|
||||
- Comprehensive readme with usage examples
|
||||
- Example scripts: basic validation, chatbot testing, benchmark regression
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,534 @@
|
||||
# Project Plan: Veritext — Semantic Text Validation Framework
|
||||
|
||||
## Overview
|
||||
|
||||
A Python library for validating text outputs against semantic criteria. Designed for
|
||||
developers building any system that produces text — chatbots, content generators,
|
||||
translation pipelines, summarisation tools — who need automated quality assurance
|
||||
beyond simple string matching.
|
||||
|
||||
**Origin story:** "I was building a feature that generated article summaries and got
|
||||
tired of manually checking if they captured the key points. Existing tools could tell
|
||||
me if two strings matched, but not if they *meant* the same thing. So I built a
|
||||
validation framework that understands semantics."
|
||||
|
||||
**Portfolio role:** A practical developer tool that demonstrates Python framework
|
||||
design, NLP evaluation techniques, and test automation integration. The project
|
||||
solves a real problem any developer working with text processing encounters.
|
||||
|
||||
**Target users:** Developers building content pipelines, chatbot teams validating
|
||||
responses, ML engineers evaluating model outputs, QA teams testing text-based features.
|
||||
|
||||
---
|
||||
|
||||
## Problem Statement
|
||||
|
||||
Text validation is hard. Traditional testing approaches fall short:
|
||||
|
||||
| Approach | Problem |
|
||||
|----------|---------|
|
||||
| Exact string match | Fails on semantically equivalent variations |
|
||||
| Substring/regex | Brittle, misses meaning entirely |
|
||||
| Manual review | Doesn't scale, subjective |
|
||||
| Generic diff tools | Show *what* changed, not *if it matters* |
|
||||
|
||||
**Example:** A summarisation system produces "The CEO announced layoffs affecting 500
|
||||
employees" one day and "500 workers will lose their jobs, the company's chief executive
|
||||
said" the next. These are semantically equivalent, but every traditional test would
|
||||
flag this as a failure.
|
||||
|
||||
Veritext answers: "Is this text output *good enough* according to my criteria?" — not
|
||||
"Is it identical?"
|
||||
|
||||
---
|
||||
|
||||
## Core Concepts
|
||||
|
||||
### Metrics (Pure Computation)
|
||||
|
||||
Metrics compute scores comparing candidate text to reference text:
|
||||
|
||||
```python
|
||||
from veritext.metrics import Bleu, Rouge
|
||||
|
||||
bleu = Bleu()
|
||||
result = bleu.score(
|
||||
candidate="The cat sat on the mat",
|
||||
reference="A cat is sitting on a mat"
|
||||
)
|
||||
# BleuResult(bleu1=0.71, bleu2=0.58, bleu3=0.45, bleu4=0.41, brevity_penalty=1.0)
|
||||
|
||||
rouge = Rouge()
|
||||
result = rouge.score(candidate, reference)
|
||||
# RougeResult(rouge1=RougeScore(...), rouge2=RougeScore(...), rouge_l=RougeScore(...))
|
||||
```
|
||||
|
||||
**Built-in metrics:**
|
||||
|
||||
| Metric | What it measures | Use case |
|
||||
|--------|------------------|----------|
|
||||
| BLEU-1 to BLEU-4 | N-gram precision | Translation, generation |
|
||||
| ROUGE-1, ROUGE-2 | N-gram recall | Summarisation |
|
||||
| ROUGE-L | Longest common subsequence | Summarisation |
|
||||
| Semantic similarity | Cosine distance of embeddings | Any meaning comparison |
|
||||
| Lexical overlap | Jaccard similarity of tokens | Simple similarity |
|
||||
| Reading level | Flesch-Kincaid grade | Accessibility |
|
||||
|
||||
**Note:** Reading level is a standalone metric (`requires_reference = False`) that analyses only the candidate text. Comparison metrics (BLEU, ROUGE, semantic) require a reference and raise `ValueError` if none provided.
|
||||
|
||||
### Validators (Decision Logic)
|
||||
|
||||
Validators wrap metrics and apply thresholds to make pass/fail decisions:
|
||||
|
||||
```python
|
||||
from veritext import validators as v
|
||||
|
||||
# Compose multiple checks
|
||||
validator = v.all_of([
|
||||
v.bleu(min_score=0.7),
|
||||
v.length(max_chars=500),
|
||||
v.readability(max_grade=8),
|
||||
])
|
||||
|
||||
from veritext.core.types import ValidationContext
|
||||
|
||||
context = ValidationContext(reference="The quick brown fox jumps over the lazy dog")
|
||||
result = validator.validate("The fast brown fox leaped over the lazy dog", context)
|
||||
# ValidationResult(passed=True, checks=[...])
|
||||
```
|
||||
|
||||
### Pytest Integration
|
||||
|
||||
Native pytest fixtures and assertions for CI/CD:
|
||||
|
||||
```python
|
||||
from veritext.pytest_plugin import validate_text
|
||||
|
||||
def test_summary_quality(summariser, document):
|
||||
summary = summariser.summarise(document)
|
||||
|
||||
validate_text(
|
||||
summary,
|
||||
reference=expected_summary,
|
||||
min_rouge=0.7,
|
||||
min_semantic=0.85,
|
||||
)
|
||||
```
|
||||
|
||||
### Regression Detection
|
||||
|
||||
Track output quality over time, catch degradations before users do:
|
||||
|
||||
```python
|
||||
from veritext.benchmark import Benchmark
|
||||
|
||||
benchmark = Benchmark("summarisation_quality", storage_path="benchmarks/")
|
||||
results = benchmark.evaluate(outputs, references, metrics=["rouge_l", "bleu4"])
|
||||
benchmark.assert_no_regression(tolerance=0.05)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Tech Stack
|
||||
|
||||
| Component | Technology | Rationale |
|
||||
|-----------|------------|-----------|
|
||||
| Core | Python 3.11+ | Target ecosystem, modern type hints |
|
||||
| Metrics | Custom implementations | Full control, understanding of algorithms |
|
||||
| Embeddings | sentence-transformers | Semantic similarity (optional) |
|
||||
| Test integration | pytest | Fixtures, plugins, assertions |
|
||||
| CLI | typer | Consistent with portfolio projects |
|
||||
| Data handling | pydantic | Validation, serialisation |
|
||||
| Storage | SQLite | Benchmark history, lightweight |
|
||||
| Output | rich | Terminal formatting |
|
||||
|
||||
---
|
||||
|
||||
## Architecture
|
||||
|
||||
### Layered Design
|
||||
|
||||
```
|
||||
┌─────────────────────────────────────────────────────┐
|
||||
│ CLI / pytest_plugin (presentation layer) │
|
||||
├─────────────────────────────────────────────────────┤
|
||||
│ validators/ (decision logic) │
|
||||
│ benchmark/ (tracking & regression) │
|
||||
├─────────────────────────────────────────────────────┤
|
||||
│ metrics/ (pure computation) │
|
||||
├─────────────────────────────────────────────────────┤
|
||||
│ core/ (shared types, tokenisation) │
|
||||
└─────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
**Dependency rule:** Each layer depends only on layers below it.
|
||||
|
||||
### Key Design Decisions
|
||||
|
||||
1. **Metrics vs Validators separation** — Metrics compute scores; validators make
|
||||
pass/fail decisions. Clear separation of concerns.
|
||||
|
||||
2. **Typed result objects** — Each metric returns a specific result type (e.g.,
|
||||
`BleuResult`, `RougeResult`), not just `float`. Full information preserved.
|
||||
|
||||
3. **Optional heavy dependencies** — `sentence-transformers` (~2GB with PyTorch) is
|
||||
optional. Core library works without ML dependencies.
|
||||
|
||||
4. **Shared tokenisation** — Single `Tokeniser` protocol used by all n-gram metrics.
|
||||
Consistent behaviour across BLEU and ROUGE.
|
||||
|
||||
5. **Explicit context** — `ValidationContext` dataclass instead of `**kwargs`.
|
||||
Type-safe, discoverable API.
|
||||
|
||||
6. **Graceful edge case handling** — Empty text returns zero scores (not errors).
|
||||
Missing reference raises clear `ValueError` for comparison metrics. Unicode
|
||||
normalised to NFC by default.
|
||||
|
||||
---
|
||||
|
||||
## Project Components
|
||||
|
||||
### Component 1: Core Module
|
||||
|
||||
Shared types, exceptions, and tokenisation.
|
||||
|
||||
**Types:**
|
||||
- `ValidationContext` — reference text and metadata for validation
|
||||
- `CheckResult` — individual check result with diagnostics
|
||||
- `ValidationResult` — aggregate result with pass/fail and all checks
|
||||
- `BatchResult` — statistics over multiple evaluations
|
||||
|
||||
**Tokeniser:**
|
||||
```python
|
||||
class Tokeniser(Protocol):
|
||||
def tokenise(self, text: str) -> list[str]: ...
|
||||
|
||||
class WordTokeniser:
|
||||
def __init__(self, lowercase: bool = True, remove_punctuation: bool = True): ...
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Component 2: Metric Engine
|
||||
|
||||
Pure implementations of text evaluation metrics.
|
||||
|
||||
**Interface:**
|
||||
```python
|
||||
class Metric(Protocol[T]):
|
||||
@property
|
||||
def name(self) -> str: ...
|
||||
|
||||
@property
|
||||
def requires_reference(self) -> bool: ...
|
||||
|
||||
def score(self, candidate: str, reference: str | list[str] | None = None) -> T:
|
||||
"""Raises ValueError if reference required but not provided."""
|
||||
...
|
||||
|
||||
def batch_score(
|
||||
self,
|
||||
candidates: list[str],
|
||||
references: list[str] | list[list[str]] | None = None,
|
||||
) -> BatchResult[T]: ...
|
||||
```
|
||||
|
||||
**Metrics:**
|
||||
- `Bleu` — BLEU-1 through BLEU-4 with brevity penalty
|
||||
- `Rouge` — ROUGE-1, ROUGE-2, ROUGE-L with precision/recall/F1
|
||||
- `Lexical` — Jaccard similarity, token overlap
|
||||
- `Readability` — Flesch-Kincaid grade level
|
||||
- `SemanticSimilarity` — Embedding cosine distance (optional dependency)
|
||||
|
||||
---
|
||||
|
||||
### Component 3: Validator Framework
|
||||
|
||||
Composable validation rules with clear pass/fail semantics.
|
||||
|
||||
**Built-in validators:**
|
||||
|
||||
| Validator | Description |
|
||||
|-----------|-------------|
|
||||
| `v.bleu(min_score, variant)` | BLEU score above minimum |
|
||||
| `v.rouge(min_score, variant)` | ROUGE score above minimum |
|
||||
| `v.semantic(min_score)` | Semantic similarity above threshold |
|
||||
| `v.length(min_chars, max_chars)` | Length constraints |
|
||||
| `v.readability(max_grade)` | Reading level constraint |
|
||||
| `v.contains(terms)` | Required terms present |
|
||||
| `v.excludes(terms)` | Forbidden terms absent |
|
||||
| `v.pattern(regex)` | Regex pattern match |
|
||||
|
||||
**Composition:**
|
||||
|
||||
```python
|
||||
# All validators must pass
|
||||
v.all_of([v.bleu(min_score=0.7), v.length(max_chars=500)])
|
||||
|
||||
# At least one must pass
|
||||
v.any_of([v.contains(["error"]), v.contains(["failed"])])
|
||||
|
||||
# Weighted scoring
|
||||
v.weighted([
|
||||
(v.bleu(min_score=0.7), 0.6),
|
||||
(v.readability(max_grade=8), 0.4),
|
||||
], min_score=0.75)
|
||||
```
|
||||
|
||||
**Reference requirements:** Validators wrapping comparison metrics (`bleu`, `rouge`, `semantic`) require `context.reference` to be set. If `None`, they raise `ValidationError` with a clear message. Constraint validators (`length`, `readability`, `contains`) do not require a reference.
|
||||
|
||||
---
|
||||
|
||||
### Component 4: Pytest Plugin
|
||||
|
||||
First-class pytest integration for CI/CD pipelines.
|
||||
|
||||
**Features:**
|
||||
- Custom assertions with detailed failure messages
|
||||
- Fixtures for common validation patterns
|
||||
- Markers for categorising text tests
|
||||
|
||||
**Usage:**
|
||||
|
||||
```python
|
||||
from veritext.pytest_plugin import validate_text
|
||||
|
||||
def test_chatbot_response():
|
||||
response = chatbot.respond("What are your hours?")
|
||||
|
||||
validate_text(
|
||||
response,
|
||||
reference="We're open Monday to Friday, 9am to 5pm.",
|
||||
min_bleu=0.6,
|
||||
min_semantic=0.8,
|
||||
max_length=500,
|
||||
)
|
||||
```
|
||||
|
||||
**Failure output:**
|
||||
|
||||
```
|
||||
FAILED test_summary.py::test_summary_quality
|
||||
AssertionError: Text failed 2 of 4 checks:
|
||||
|
||||
✗ rouge: 0.58 (minimum: 0.70)
|
||||
✗ semantic: 0.72 (minimum: 0.85)
|
||||
✓ length: 342 (maximum: 500)
|
||||
✓ readability: 6.2 (maximum: 8)
|
||||
|
||||
Candidate: "The company reported losses..."
|
||||
Reference: "Financial results showed significant decline..."
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Component 5: Benchmark & Regression Detection
|
||||
|
||||
Track quality over time, catch degradations automatically.
|
||||
|
||||
**Features:**
|
||||
- Store historical metric values in SQLite
|
||||
- Statistical regression detection
|
||||
- Configurable tolerance thresholds
|
||||
- CI integration for blocking degradations
|
||||
|
||||
**Usage:**
|
||||
|
||||
```python
|
||||
from veritext.benchmark import Benchmark
|
||||
|
||||
benchmark = Benchmark("chatbot_quality", storage_path="benchmarks/")
|
||||
|
||||
# Record current run (returns BenchmarkRun with metrics and metadata)
|
||||
run = benchmark.evaluate(
|
||||
candidates=current_outputs,
|
||||
references=expected_outputs,
|
||||
metrics=["rouge_l", "semantic"]
|
||||
)
|
||||
# run.metrics = {"rouge_l": 0.82, "semantic": 0.89}
|
||||
|
||||
# Compare against historical baseline
|
||||
regression = benchmark.check_regression(tolerance=0.05, window=10)
|
||||
|
||||
if regression.detected:
|
||||
print(f"Quality dropped: {regression.summary}")
|
||||
|
||||
# In CI: fail the build on regression
|
||||
benchmark.assert_no_regression(tolerance=0.05)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Component 6: CLI Tool
|
||||
|
||||
Command-line interface for quick validation and benchmarking.
|
||||
|
||||
```bash
|
||||
# Validate a single text
|
||||
$ veritext validate "Your text here" --reference "Expected text" --metrics bleu,rouge
|
||||
|
||||
# Validate from files
|
||||
$ veritext validate --file outputs.jsonl --reference-file expected.jsonl
|
||||
|
||||
# Run benchmark
|
||||
$ veritext benchmark run summarisation --inputs docs/ --references refs/
|
||||
|
||||
# Show benchmark history
|
||||
$ veritext benchmark show summarisation --last 20
|
||||
|
||||
# Check for regression
|
||||
$ veritext benchmark check summarisation --tolerance 0.05
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Example Use Cases
|
||||
|
||||
### Use Case 1: Chatbot Response Validation
|
||||
|
||||
```python
|
||||
from veritext import validators as v
|
||||
from veritext.core.types import ValidationContext
|
||||
|
||||
# Define acceptable response criteria
|
||||
response_validator = v.all_of([
|
||||
v.length(max_chars=500),
|
||||
v.readability(max_grade=8),
|
||||
v.excludes(terms=["I don't know", "I'm not sure"]),
|
||||
])
|
||||
|
||||
def test_chatbot_responds_helpfully():
|
||||
response = chatbot.respond("How do I reset my password?")
|
||||
context = ValidationContext()
|
||||
result = response_validator.validate(response, context)
|
||||
assert result.passed, result.failure_summary
|
||||
```
|
||||
|
||||
### Use Case 2: Summarisation Quality Gate
|
||||
|
||||
```python
|
||||
from veritext.pytest_plugin import validate_text
|
||||
|
||||
def test_summary_captures_key_points():
|
||||
article = load_article("financial_report.txt")
|
||||
summary = summariser.summarise(article)
|
||||
|
||||
validate_text(
|
||||
summary,
|
||||
reference=load_reference_summary("financial_report_summary.txt"),
|
||||
min_rouge=0.65,
|
||||
min_semantic=0.80,
|
||||
max_length=300,
|
||||
)
|
||||
```
|
||||
|
||||
### Use Case 3: Translation Quality Monitoring
|
||||
|
||||
```python
|
||||
from veritext.benchmark import Benchmark
|
||||
|
||||
benchmark = Benchmark("translation_en_de", storage_path="benchmarks/")
|
||||
|
||||
# Nightly CI job
|
||||
results = benchmark.evaluate(
|
||||
candidates=translate_batch(test_documents),
|
||||
references=human_translations,
|
||||
metrics=["bleu4", "semantic"]
|
||||
)
|
||||
|
||||
# Block deployment if quality drops
|
||||
benchmark.assert_no_regression(tolerance=0.03)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Success Criteria
|
||||
|
||||
- [ ] BLEU/ROUGE implementations match reference implementations (nltk, rouge-score)
|
||||
- [ ] Semantic similarity correlates with human judgement on test pairs
|
||||
- [ ] Pytest plugin installs cleanly via `uv pip install veritext`
|
||||
- [ ] Validation of 1000 text pairs completes in <5 seconds (excluding embeddings)
|
||||
- [ ] Benchmark regression detection has <5% false positive rate
|
||||
- [ ] Edge cases handled gracefully (empty text, None reference, Unicode)
|
||||
- [ ] Documentation includes working examples for each use case
|
||||
- [ ] All code passes ruff, mypy strict, and pytest with ≥80% coverage
|
||||
- [ ] Can explain design decisions and metric theory in interview
|
||||
|
||||
---
|
||||
|
||||
## Skills Demonstrated
|
||||
|
||||
| Skill | How Veritext demonstrates it |
|
||||
|-------|------------------------------|
|
||||
| Python framework design | Composable validators, clean API, plugin architecture |
|
||||
| Test automation | Native pytest integration, CI/CD workflows |
|
||||
| NLP evaluation metrics | BLEU, ROUGE, semantic similarity implementations |
|
||||
| Data analysis | Statistical regression detection, batch processing |
|
||||
| CLI development | Typer-based interface, rich output |
|
||||
| Software architecture | Layered design, clear separation of concerns |
|
||||
| Documentation | Comprehensive readme, examples |
|
||||
| Quality engineering | High test coverage, type safety, linting |
|
||||
|
||||
---
|
||||
|
||||
## What Makes This Project Credible
|
||||
|
||||
1. **Solves a real problem** — Anyone building text-based features faces validation
|
||||
challenges.
|
||||
|
||||
2. **Not tied to a specific technology** — Works with any text source (chatbots, LLMs,
|
||||
translation APIs, content generators). It's a general-purpose tool, not an "LLM
|
||||
testing framework."
|
||||
|
||||
3. **Practical scope** — Not trying to reinvent pytest or build an ML platform. Focused
|
||||
on one thing: validating text quality.
|
||||
|
||||
4. **Demonstrates depth** — Implementing BLEU/ROUGE from understanding (not just
|
||||
wrapping libraries) shows knowledge of how these metrics work.
|
||||
|
||||
5. **Natural portfolio narrative** — "I was building X and needed a better way to test
|
||||
it, so I built this tool." Every interviewer has faced similar problems.
|
||||
|
||||
---
|
||||
|
||||
## Portfolio Demos (Future)
|
||||
|
||||
Interactive demos to showcase Veritext without requiring installation.
|
||||
|
||||
### Streamlit Demo
|
||||
|
||||
A quick interactive web UI for general visitors and recruiters.
|
||||
|
||||
**Features:**
|
||||
- Text input boxes (candidate + reference)
|
||||
- Metric selector (BLEU, ROUGE, lexical, readability)
|
||||
- Threshold sliders for pass/fail validation
|
||||
- Results table with scores and status
|
||||
|
||||
**Deployment:** Self-hosted on homeserver (e.g., `veritext.kschappell.com`)
|
||||
|
||||
**Effort:** ~30 minutes
|
||||
|
||||
### Jupyter Notebook Collection
|
||||
|
||||
Deep-dive notebooks targeting data science and ML recruiters.
|
||||
|
||||
**Notebooks:**
|
||||
|
||||
| Notebook | Purpose |
|
||||
|----------|---------|
|
||||
| `01-metrics-overview.ipynb` | Introduction to each metric with visualisations |
|
||||
| `02-batch-evaluation.ipynb` | Evaluating model outputs at scale, statistical analysis |
|
||||
| `03-regression-detection.ipynb` | Tracking quality over time, detecting degradation |
|
||||
| `04-chatbot-validation.ipynb` | Real-world use case: validating chatbot responses |
|
||||
|
||||
**Hosting:** JupyterLite (static files, runs in browser via WebAssembly)
|
||||
|
||||
**Deployment:** Self-hosted alongside Streamlit demo
|
||||
|
||||
**Why both:**
|
||||
|
||||
| Demo Type | Audience | Value |
|
||||
|-----------|----------|-------|
|
||||
| Streamlit | General visitors | Quick, interactive, no friction |
|
||||
| Notebooks | Data/ML recruiters | Shows analytical depth, speaks their language |
|
||||
@@ -34,11 +34,13 @@ CHATBOT_RESPONSES = {
|
||||
# Fixtures for common test setup
|
||||
@pytest.fixture
|
||||
def greeting_response() -> str:
|
||||
"""Provide a sample greeting response."""
|
||||
return CHATBOT_RESPONSES["greeting"]["response"]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def weather_response() -> str:
|
||||
"""Provide a sample weather response."""
|
||||
return CHATBOT_RESPONSES["weather"]["response"]
|
||||
|
||||
|
||||
@@ -47,6 +49,7 @@ class TestResponseQuality:
|
||||
"""Test chatbot response quality using Veritext."""
|
||||
|
||||
def test_greeting_length(self, greeting_response: str) -> None:
|
||||
"""Greeting responses should be concise."""
|
||||
validate_text(
|
||||
greeting_response,
|
||||
min_length=10,
|
||||
@@ -54,12 +57,14 @@ class TestResponseQuality:
|
||||
)
|
||||
|
||||
def test_greeting_readability(self, greeting_response: str) -> None:
|
||||
"""Greeting responses should be easy to read."""
|
||||
validate_text(
|
||||
greeting_response,
|
||||
max_reading_grade=8.0,
|
||||
)
|
||||
|
||||
def test_greeting_contains_keywords(self, greeting_response: str) -> None:
|
||||
"""Greeting should contain expected terms."""
|
||||
validate_text(
|
||||
greeting_response,
|
||||
must_contain=["help"],
|
||||
|
||||
@@ -20,8 +20,10 @@ def compute_baseline(
|
||||
if not runs:
|
||||
return {}
|
||||
|
||||
# Take up to `window` runs
|
||||
recent_runs = runs[:window]
|
||||
|
||||
# Collect all metric values
|
||||
metric_values: dict[str, list[float]] = {}
|
||||
for run in recent_runs:
|
||||
for metric_name, value in run.metrics.items():
|
||||
@@ -29,6 +31,7 @@ def compute_baseline(
|
||||
metric_values[metric_name] = []
|
||||
metric_values[metric_name].append(value)
|
||||
|
||||
# Compute averages
|
||||
return {
|
||||
metric: sum(values) / len(values) for metric, values in metric_values.items()
|
||||
}
|
||||
@@ -54,6 +57,7 @@ def detect_regression(
|
||||
RegressionReport with comparison results.
|
||||
"""
|
||||
if not baseline:
|
||||
# No baseline means no regression possible
|
||||
return RegressionReport(
|
||||
detected=False,
|
||||
baseline=baseline,
|
||||
@@ -70,6 +74,7 @@ def detect_regression(
|
||||
delta = current_value - baseline_value
|
||||
deltas[metric] = delta
|
||||
|
||||
# Check if this metric regressed beyond tolerance
|
||||
if delta < -tolerance:
|
||||
detected = True
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ from veritext.core.exceptions import RegressionDetectedError
|
||||
from veritext.metrics.bleu import Bleu
|
||||
from veritext.metrics.rouge import Rouge
|
||||
|
||||
# Default metrics to use for evaluation
|
||||
DEFAULT_METRICS = ["rouge_l", "bleu4"]
|
||||
|
||||
|
||||
@@ -35,11 +36,13 @@ class Benchmark:
|
||||
self._storage_path = Path(storage_path)
|
||||
self._storage = BenchmarkStorage(self._storage_path / f"{name}.db")
|
||||
|
||||
# Initialise metrics
|
||||
self._bleu = Bleu()
|
||||
self._rouge = Rouge()
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the benchmark name."""
|
||||
return self._name
|
||||
|
||||
def _compute_metrics(
|
||||
@@ -48,6 +51,7 @@ class Benchmark:
|
||||
references: list[str] | list[list[str]],
|
||||
metric_names: list[str],
|
||||
) -> dict[str, float]:
|
||||
"""Compute requested metrics for the given samples."""
|
||||
results: dict[str, float] = {}
|
||||
|
||||
for metric_name in metric_names:
|
||||
@@ -66,6 +70,7 @@ class Benchmark:
|
||||
"rouge_l_fmeasure",
|
||||
):
|
||||
rouge_result = self._rouge.batch_score(candidates, references)
|
||||
# Map short names to stat names
|
||||
stat_name = metric_name
|
||||
if metric_name == "rouge1":
|
||||
stat_name = "rouge1_fmeasure"
|
||||
@@ -133,6 +138,7 @@ class Benchmark:
|
||||
runs = self._storage.get_runs(self._name)
|
||||
|
||||
if not runs:
|
||||
# No runs at all
|
||||
return RegressionReport(
|
||||
detected=False,
|
||||
baseline={},
|
||||
@@ -142,6 +148,7 @@ class Benchmark:
|
||||
)
|
||||
|
||||
current_run = runs[0]
|
||||
# Baseline excludes the current run
|
||||
historical_runs = runs[1:]
|
||||
baseline = compute_baseline(historical_runs, window=window)
|
||||
|
||||
|
||||
@@ -20,7 +20,23 @@ class BenchmarkStorage:
|
||||
db_path: Path to the SQLite database file.
|
||||
"""
|
||||
self._db_path = db_path
|
||||
self._ensure_parent_exists()
|
||||
self._init_database()
|
||||
|
||||
def _ensure_parent_exists(self) -> None:
|
||||
"""Ensure the parent directory exists."""
|
||||
self._db_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def _get_connection(self) -> sqlite3.Connection:
|
||||
"""Get a database connection with WAL mode enabled."""
|
||||
conn = sqlite3.connect(str(self._db_path), timeout=30.0)
|
||||
conn.execute("PRAGMA journal_mode=WAL")
|
||||
conn.execute("PRAGMA foreign_keys=ON")
|
||||
conn.row_factory = sqlite3.Row
|
||||
return conn
|
||||
|
||||
def _init_database(self) -> None:
|
||||
"""Create tables if they don't exist."""
|
||||
try:
|
||||
with self._get_connection() as conn:
|
||||
conn.executescript("""
|
||||
@@ -46,13 +62,6 @@ class BenchmarkStorage:
|
||||
except sqlite3.Error as e:
|
||||
raise StorageError(f"Failed to initialise database: {e}") from e
|
||||
|
||||
def _get_connection(self) -> sqlite3.Connection:
|
||||
conn = sqlite3.connect(str(self._db_path), timeout=30.0)
|
||||
conn.execute("PRAGMA journal_mode=WAL")
|
||||
conn.execute("PRAGMA foreign_keys=ON")
|
||||
conn.row_factory = sqlite3.Row
|
||||
return conn
|
||||
|
||||
def save_run(self, run: BenchmarkRun) -> None:
|
||||
"""
|
||||
Persist a benchmark run.
|
||||
@@ -65,6 +74,7 @@ class BenchmarkStorage:
|
||||
"""
|
||||
try:
|
||||
with self._get_connection() as conn:
|
||||
# Insert the run
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO benchmark_runs
|
||||
@@ -81,6 +91,7 @@ class BenchmarkStorage:
|
||||
),
|
||||
)
|
||||
|
||||
# Insert metrics
|
||||
for metric_name, value in run.metrics.items():
|
||||
conn.execute(
|
||||
"""
|
||||
@@ -129,6 +140,7 @@ class BenchmarkStorage:
|
||||
|
||||
runs = []
|
||||
for row in rows:
|
||||
# Get metrics for this run
|
||||
metrics_rows = conn.execute(
|
||||
"SELECT metric_name, value FROM benchmark_metrics WHERE run_id = ?",
|
||||
(row["id"],),
|
||||
@@ -154,5 +166,14 @@ class BenchmarkStorage:
|
||||
raise StorageError(f"Failed to retrieve benchmark runs: {e}") from e
|
||||
|
||||
def get_latest_run(self, benchmark_name: str) -> BenchmarkRun | None:
|
||||
"""
|
||||
Get the most recent run for a benchmark.
|
||||
|
||||
Args:
|
||||
benchmark_name: Name of the benchmark.
|
||||
|
||||
Returns:
|
||||
The most recent BenchmarkRun, or None if no runs exist.
|
||||
"""
|
||||
runs = self.get_runs(benchmark_name, limit=1)
|
||||
return runs[0] if runs else None
|
||||
|
||||
@@ -45,10 +45,28 @@ def format_validation_table(
|
||||
|
||||
|
||||
def format_validation_json(results: dict[str, float]) -> str:
|
||||
"""
|
||||
Format validation results as JSON.
|
||||
|
||||
Args:
|
||||
results: Dictionary of metric names to scores.
|
||||
|
||||
Returns:
|
||||
JSON string.
|
||||
"""
|
||||
return json.dumps(results, indent=2)
|
||||
|
||||
|
||||
def format_validation_simple(results: dict[str, float]) -> str:
|
||||
"""
|
||||
Format validation results as simple text output.
|
||||
|
||||
Args:
|
||||
results: Dictionary of metric names to scores.
|
||||
|
||||
Returns:
|
||||
Simple text string with one metric per line.
|
||||
"""
|
||||
lines = [f"{metric}: {score:.4f}" for metric, score in sorted(results.items())]
|
||||
return "\n".join(lines)
|
||||
|
||||
@@ -68,6 +86,7 @@ def format_benchmark_history(runs: list[BenchmarkRun]) -> Table:
|
||||
table.add_column("No runs found")
|
||||
return table
|
||||
|
||||
# Get all metric names from the runs
|
||||
metric_names: set[str] = set()
|
||||
for run in runs:
|
||||
metric_names.update(run.metrics.keys())
|
||||
@@ -104,6 +123,7 @@ def format_regression_report(report: RegressionReport) -> Panel:
|
||||
)
|
||||
return Panel(content, title="Regression Check", border_style="green")
|
||||
|
||||
# Build regression details
|
||||
lines = [
|
||||
"[red]Regression detected![/red]",
|
||||
f"Tolerance: {report.tolerance:.2%}",
|
||||
|
||||
@@ -7,6 +7,7 @@ from pathlib import Path
|
||||
|
||||
@dataclass
|
||||
class TextPair:
|
||||
"""A candidate-reference text pair for validation."""
|
||||
|
||||
candidate: str
|
||||
reference: str
|
||||
@@ -91,6 +92,7 @@ def read_paired_jsonl(candidates_path: Path, references_path: Path) -> list[Text
|
||||
|
||||
|
||||
def _read_text_jsonl(path: Path, label: str) -> list[str]:
|
||||
"""Read text values from a JSONL file with 'text' key per line."""
|
||||
if not path.exists():
|
||||
raise FileNotFoundError(f"{label.capitalize()} file not found: {path}")
|
||||
|
||||
|
||||
@@ -11,16 +11,19 @@ from veritext.metrics.bleu import Bleu
|
||||
from veritext.metrics.lexical import Lexical
|
||||
from veritext.metrics.rouge import Rouge
|
||||
|
||||
# Available metrics
|
||||
AVAILABLE_METRICS = frozenset(
|
||||
{"bleu", "bleu1", "bleu2", "bleu3", "bleu4", "rouge", "rouge_l", "lexical"}
|
||||
)
|
||||
|
||||
# Lazily-initialised metric instances
|
||||
_bleu: Bleu | None = None
|
||||
_rouge: Rouge | None = None
|
||||
_lexical: Lexical | None = None
|
||||
|
||||
|
||||
def _get_bleu() -> Bleu:
|
||||
"""Get or create the BLEU metric instance."""
|
||||
global _bleu
|
||||
if _bleu is None:
|
||||
_bleu = Bleu()
|
||||
@@ -28,6 +31,7 @@ def _get_bleu() -> Bleu:
|
||||
|
||||
|
||||
def _get_rouge() -> Rouge:
|
||||
"""Get or create the ROUGE metric instance."""
|
||||
global _rouge
|
||||
if _rouge is None:
|
||||
_rouge = Rouge()
|
||||
@@ -35,11 +39,19 @@ def _get_rouge() -> Rouge:
|
||||
|
||||
|
||||
def _get_lexical() -> Lexical:
|
||||
"""Get or create the lexical metric instance."""
|
||||
global _lexical
|
||||
if _lexical is None:
|
||||
_lexical = Lexical()
|
||||
return _lexical
|
||||
|
||||
|
||||
# Metric registry: maps metric names to (result_keys, single_extractor, batch_extractor)
|
||||
# - result_keys: output keys to populate
|
||||
# - single_extractor: function(candidate, reference) -> dict of results
|
||||
# - batch_extractor: function(candidates, references) -> dict of results
|
||||
def _bleu_single(candidate: str, reference: str, key: str) -> dict[str, float]:
|
||||
"""Extract a BLEU score for single mode."""
|
||||
result = _get_bleu().score(candidate, reference)
|
||||
return {key: getattr(result, key)}
|
||||
|
||||
@@ -47,28 +59,33 @@ def _bleu_single(candidate: str, reference: str, key: str) -> dict[str, float]:
|
||||
def _bleu_batch(
|
||||
candidates: list[str], references: list[str], key: str
|
||||
) -> dict[str, float]:
|
||||
"""Extract a BLEU score for batch mode."""
|
||||
batch = _get_bleu().batch_score(candidates, references)
|
||||
stats = batch.stats.get(key)
|
||||
return {key: stats.mean} if stats else {}
|
||||
|
||||
|
||||
def _rouge_single(candidate: str, reference: str) -> dict[str, float]:
|
||||
"""Extract ROUGE-L F-measure for single mode."""
|
||||
result = _get_rouge().score(candidate, reference)
|
||||
return {"rouge_l": result.rouge_l.fmeasure}
|
||||
|
||||
|
||||
def _rouge_batch(candidates: list[str], references: list[str]) -> dict[str, float]:
|
||||
"""Extract ROUGE-L F-measure for batch mode."""
|
||||
batch = _get_rouge().batch_score(candidates, references)
|
||||
stats = batch.stats.get("rouge_l_fmeasure")
|
||||
return {"rouge_l": stats.mean} if stats else {}
|
||||
|
||||
|
||||
def _lexical_single(candidate: str, reference: str) -> dict[str, float]:
|
||||
"""Extract lexical scores for single mode."""
|
||||
result = _get_lexical().score(candidate, reference)
|
||||
return {"jaccard": result.jaccard, "token_overlap": result.token_overlap}
|
||||
|
||||
|
||||
def _lexical_batch(candidates: list[str], references: list[str]) -> dict[str, float]:
|
||||
"""Extract lexical scores for batch mode."""
|
||||
batch = _get_lexical().batch_score(candidates, references)
|
||||
results: dict[str, float] = {}
|
||||
jaccard_stats = batch.stats.get("jaccard")
|
||||
@@ -85,6 +102,7 @@ def _compute_metrics(
|
||||
reference: str,
|
||||
metric_names: list[str],
|
||||
) -> dict[str, float]:
|
||||
"""Compute requested metrics for a single text pair."""
|
||||
results: dict[str, float] = {}
|
||||
|
||||
for metric in metric_names:
|
||||
@@ -105,6 +123,7 @@ def _compute_batch_metrics(
|
||||
references: list[str],
|
||||
metric_names: list[str],
|
||||
) -> dict[str, float]:
|
||||
"""Compute average metrics for a batch of text pairs."""
|
||||
results: dict[str, float] = {}
|
||||
|
||||
for metric in metric_names:
|
||||
@@ -121,7 +140,10 @@ def _compute_batch_metrics(
|
||||
|
||||
|
||||
def _parse_metrics(metrics_str: str) -> list[str]:
|
||||
"""Parse comma-separated metric names."""
|
||||
metrics = [m.strip().lower() for m in metrics_str.split(",")]
|
||||
|
||||
# Validate metric names
|
||||
invalid = [m for m in metrics if m not in AVAILABLE_METRICS]
|
||||
if invalid:
|
||||
raise typer.BadParameter(
|
||||
@@ -179,17 +201,21 @@ def validate(
|
||||
Use file mode for batches:
|
||||
veritext validate -f outputs.jsonl -m bleu,rouge
|
||||
"""
|
||||
# Parse and validate metric names
|
||||
try:
|
||||
metric_names = _parse_metrics(metrics)
|
||||
except typer.BadParameter as e:
|
||||
console.print(f"[red]Error:[/red] {e}")
|
||||
raise typer.Exit(code=1) from e
|
||||
|
||||
# Validate output format
|
||||
if output not in ("table", "json", "simple"):
|
||||
console.print(f"[red]Error:[/red] Invalid output format: {output}")
|
||||
raise typer.Exit(code=1)
|
||||
|
||||
# Determine mode: inline vs file
|
||||
if file is not None:
|
||||
# File mode
|
||||
try:
|
||||
if reference_file is not None:
|
||||
pairs = read_paired_jsonl(file, reference_file)
|
||||
@@ -210,9 +236,11 @@ def validate(
|
||||
console.print(f"[dim]Evaluated {len(pairs)} text pairs.[/dim]\n")
|
||||
|
||||
elif text is not None and reference is not None:
|
||||
# Inline mode
|
||||
results = _compute_metrics(text, reference, metric_names)
|
||||
|
||||
else:
|
||||
# Invalid usage
|
||||
console.print(
|
||||
"[red]Error:[/red] Provide either text and --reference, "
|
||||
"or --file for batch mode."
|
||||
|
||||
@@ -57,4 +57,5 @@ class VeritextSettings(BaseSettings):
|
||||
|
||||
@lru_cache
|
||||
def get_settings() -> VeritextSettings:
|
||||
"""Get the cached settings instance."""
|
||||
return VeritextSettings()
|
||||
|
||||
@@ -24,12 +24,14 @@ def configure_logging(
|
||||
level = level or settings.log_level
|
||||
log_format = log_format or settings.log_format
|
||||
|
||||
# Configure standard library logging
|
||||
logging.basicConfig(
|
||||
format="%(message)s",
|
||||
stream=sys.stderr,
|
||||
level=getattr(logging, level),
|
||||
)
|
||||
|
||||
# Shared processors
|
||||
shared_processors: list[Any] = [
|
||||
structlog.contextvars.merge_contextvars,
|
||||
structlog.processors.add_log_level,
|
||||
@@ -40,12 +42,14 @@ def configure_logging(
|
||||
]
|
||||
|
||||
if log_format == "json":
|
||||
# JSON output for production/log aggregation
|
||||
processors = [
|
||||
*shared_processors,
|
||||
structlog.processors.format_exc_info,
|
||||
structlog.processors.JSONRenderer(),
|
||||
]
|
||||
else:
|
||||
# Console output for development
|
||||
processors = [
|
||||
*shared_processors,
|
||||
structlog.dev.ConsoleRenderer(colors=True),
|
||||
@@ -63,4 +67,13 @@ def configure_logging(
|
||||
|
||||
|
||||
def get_logger(name: str | None = None) -> structlog.stdlib.BoundLogger:
|
||||
"""
|
||||
Get a logger instance.
|
||||
|
||||
Args:
|
||||
name: Logger name. Uses 'veritext' if not provided.
|
||||
|
||||
Returns:
|
||||
A bound logger instance.
|
||||
"""
|
||||
return structlog.get_logger(name or "veritext")
|
||||
|
||||
@@ -10,7 +10,17 @@ NormalisationForm = Literal["NFC", "NFD", "NFKC", "NFKD"]
|
||||
class Tokeniser(Protocol):
|
||||
"""Protocol for text tokenisers."""
|
||||
|
||||
def tokenise(self, text: str) -> list[str]: ...
|
||||
def tokenise(self, text: str) -> list[str]:
|
||||
"""
|
||||
Tokenise text into a list of tokens.
|
||||
|
||||
Args:
|
||||
text: The text to tokenise.
|
||||
|
||||
Returns:
|
||||
List of tokens.
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
class WordTokeniser:
|
||||
@@ -38,6 +48,7 @@ class WordTokeniser:
|
||||
self.remove_punctuation = remove_punctuation
|
||||
self.normalisation_form: NormalisationForm = normalisation_form
|
||||
|
||||
# Pattern for punctuation removal (keeps alphanumeric and Unicode letters)
|
||||
self._punctuation_pattern = re.compile(r"[^\w\s]", re.UNICODE)
|
||||
|
||||
def tokenise(self, text: str) -> list[str]:
|
||||
@@ -53,14 +64,18 @@ class WordTokeniser:
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
# Unicode normalisation
|
||||
normalised = unicodedata.normalize(self.normalisation_form, text)
|
||||
|
||||
# Lowercase if requested
|
||||
if self.lowercase:
|
||||
normalised = normalised.lower()
|
||||
|
||||
# Remove punctuation if requested
|
||||
if self.remove_punctuation:
|
||||
normalised = self._punctuation_pattern.sub(" ", normalised)
|
||||
|
||||
# Split on whitespace and filter empty strings
|
||||
tokens = normalised.split()
|
||||
|
||||
return tokens
|
||||
|
||||
@@ -51,10 +51,12 @@ class ValidationResult(BaseModel):
|
||||
|
||||
@property
|
||||
def failed_checks(self) -> list[CheckResult]:
|
||||
"""Return only the checks that failed."""
|
||||
return [check for check in self.checks if not check.passed]
|
||||
|
||||
@property
|
||||
def failure_summary(self) -> str:
|
||||
"""Return a summary of all failed checks."""
|
||||
failed = self.failed_checks
|
||||
if not failed:
|
||||
return "All checks passed."
|
||||
|
||||
@@ -61,6 +61,7 @@ class AggregateStats(BaseModel):
|
||||
sorted_values = sorted(values)
|
||||
|
||||
def percentile(p: int) -> float:
|
||||
"""Compute percentile using linear interpolation."""
|
||||
if n == 1:
|
||||
return sorted_values[0]
|
||||
k = (n - 1) * p / 100
|
||||
@@ -99,10 +100,12 @@ class Metric(Protocol[T]):
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this metric."""
|
||||
...
|
||||
|
||||
@property
|
||||
def requires_reference(self) -> bool:
|
||||
"""Return whether this metric requires reference text."""
|
||||
...
|
||||
|
||||
def score(self, candidate: str, reference: str | list[str] | None = None) -> T:
|
||||
|
||||
@@ -7,10 +7,9 @@ from veritext.core.tokenisation import WordTokeniser
|
||||
from veritext.metrics.base import AggregateStats, BatchResult
|
||||
from veritext.metrics.results import BleuResult
|
||||
|
||||
MAX_NGRAM_ORDER = 4
|
||||
|
||||
|
||||
def _get_ngrams(tokens: list[str], n: int) -> Counter[tuple[str, ...]]:
|
||||
"""Extract n-grams from a list of tokens."""
|
||||
if n > len(tokens):
|
||||
return Counter()
|
||||
return Counter(tuple(tokens[i : i + n]) for i in range(len(tokens) - n + 1))
|
||||
@@ -31,12 +30,14 @@ def _modified_precision(
|
||||
if not candidate_ngrams:
|
||||
return 0, 0
|
||||
|
||||
# Get max count for each n-gram across all references
|
||||
max_ref_counts: Counter[tuple[str, ...]] = Counter()
|
||||
for ref_tokens in reference_token_lists:
|
||||
ref_ngrams = _get_ngrams(ref_tokens, n)
|
||||
for ngram, count in ref_ngrams.items():
|
||||
max_ref_counts[ngram] = max(max_ref_counts[ngram], count)
|
||||
|
||||
# Clip candidate counts to max reference counts
|
||||
clipped_count = 0
|
||||
for ngram, count in candidate_ngrams.items():
|
||||
clipped_count += min(count, max_ref_counts[ngram])
|
||||
@@ -53,6 +54,7 @@ def _brevity_penalty(candidate_len: int, reference_lens: list[int]) -> float:
|
||||
if candidate_len == 0:
|
||||
return 0.0
|
||||
|
||||
# Find closest reference length
|
||||
closest_ref_len = min(reference_lens, key=lambda r: (abs(r - candidate_len), r))
|
||||
|
||||
if candidate_len >= closest_ref_len:
|
||||
@@ -80,10 +82,12 @@ class Bleu:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this metric."""
|
||||
return "bleu"
|
||||
|
||||
@property
|
||||
def requires_reference(self) -> bool:
|
||||
"""Return whether this metric requires reference text."""
|
||||
return True
|
||||
|
||||
def score(
|
||||
@@ -105,21 +109,26 @@ class Bleu:
|
||||
if reference is None:
|
||||
raise ValueError("BLEU requires reference text")
|
||||
|
||||
# Normalise reference to list
|
||||
references = [reference] if isinstance(reference, str) else reference
|
||||
|
||||
# Tokenise
|
||||
candidate_tokens = self._tokeniser.tokenise(candidate)
|
||||
reference_token_lists = [self._tokeniser.tokenise(r) for r in references]
|
||||
|
||||
# Handle empty candidate
|
||||
if not candidate_tokens:
|
||||
return BleuResult(
|
||||
bleu1=0.0, bleu2=0.0, bleu3=0.0, bleu4=0.0, brevity_penalty=0.0
|
||||
)
|
||||
|
||||
# Handle empty references
|
||||
if all(not ref for ref in reference_token_lists):
|
||||
raise ValueError("Reference text cannot be empty")
|
||||
|
||||
# Compute modified precisions for n=1,2,3,4
|
||||
precisions = []
|
||||
for n in range(1, MAX_NGRAM_ORDER + 1):
|
||||
for n in range(1, 5):
|
||||
clipped, total = _modified_precision(
|
||||
candidate_tokens, reference_token_lists, n
|
||||
)
|
||||
@@ -128,19 +137,22 @@ class Bleu:
|
||||
else:
|
||||
precisions.append(clipped / total)
|
||||
|
||||
# Compute BLEU scores using geometric mean
|
||||
bp = _brevity_penalty(
|
||||
len(candidate_tokens),
|
||||
[len(ref) for ref in reference_token_lists],
|
||||
)
|
||||
|
||||
def geometric_mean(p_list: list[float]) -> float:
|
||||
"""Compute geometric mean with smoothing for zeros."""
|
||||
if any(p == 0.0 for p in p_list):
|
||||
return 0.0
|
||||
log_sum = sum(math.log(p) for p in p_list)
|
||||
return math.exp(log_sum / len(p_list))
|
||||
|
||||
bleu_scores = []
|
||||
for n in range(1, MAX_NGRAM_ORDER + 1):
|
||||
for n in range(1, 5):
|
||||
# BLEU-n uses precisions 1 through n
|
||||
bleu_n = bp * geometric_mean(precisions[:n])
|
||||
bleu_scores.append(bleu_n)
|
||||
|
||||
@@ -184,6 +196,7 @@ class Bleu:
|
||||
ref: str | list[str] = references[i]
|
||||
results.append(self.score(cand, ref))
|
||||
|
||||
# Compute aggregate statistics for each score type
|
||||
stats = {
|
||||
"bleu1": AggregateStats.from_values([r.bleu1 for r in results]),
|
||||
"bleu2": AggregateStats.from_values([r.bleu2 for r in results]),
|
||||
|
||||
@@ -13,14 +13,22 @@ class Lexical:
|
||||
"""
|
||||
|
||||
def __init__(self, tokeniser: WordTokeniser | None = None) -> None:
|
||||
"""
|
||||
Initialise the lexical metric.
|
||||
|
||||
Args:
|
||||
tokeniser: Tokeniser to use. Defaults to WordTokeniser().
|
||||
"""
|
||||
self._tokeniser = tokeniser or WordTokeniser()
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this metric."""
|
||||
return "lexical"
|
||||
|
||||
@property
|
||||
def requires_reference(self) -> bool:
|
||||
"""Return whether this metric requires reference text."""
|
||||
return True
|
||||
|
||||
def score(
|
||||
@@ -43,14 +51,18 @@ class Lexical:
|
||||
if reference is None:
|
||||
raise ValueError("Lexical similarity requires reference text")
|
||||
|
||||
# Use first reference if list provided
|
||||
ref_text = reference[0] if isinstance(reference, list) else reference
|
||||
|
||||
# Tokenise
|
||||
candidate_tokens = self._tokeniser.tokenise(candidate)
|
||||
reference_tokens = self._tokeniser.tokenise(ref_text)
|
||||
|
||||
# Handle empty reference
|
||||
if not reference_tokens:
|
||||
raise ValueError("Reference text cannot be empty")
|
||||
|
||||
# Handle empty candidate
|
||||
if not candidate_tokens:
|
||||
return LexicalResult(jaccard=0.0, token_overlap=0.0)
|
||||
|
||||
@@ -97,6 +109,7 @@ class Lexical:
|
||||
ref: str | list[str] = references[i]
|
||||
results.append(self.score(cand, ref))
|
||||
|
||||
# Compute aggregate statistics
|
||||
stats = {
|
||||
"jaccard": AggregateStats.from_values([r.jaccard for r in results]),
|
||||
"token_overlap": AggregateStats.from_values(
|
||||
|
||||
@@ -5,19 +5,25 @@ import re
|
||||
from veritext.metrics.base import AggregateStats, BatchResult
|
||||
from veritext.metrics.results import ReadabilityResult
|
||||
|
||||
# Sentence-ending punctuation pattern
|
||||
_SENTENCE_ENDINGS = re.compile(r"[.!?]+")
|
||||
|
||||
# Vowel pattern for syllable counting
|
||||
_VOWELS = re.compile(r"[aeiouy]+", re.IGNORECASE)
|
||||
|
||||
FK_GRADE_WORDS_PER_SENTENCE = 0.39
|
||||
FK_GRADE_SYLLABLES_PER_WORD = 11.8
|
||||
FK_GRADE_CONSTANT = 15.59
|
||||
|
||||
FRE_CONSTANT = 206.835
|
||||
FRE_WORDS_PER_SENTENCE = 1.015
|
||||
FRE_SYLLABLES_PER_WORD = 84.6
|
||||
|
||||
|
||||
def _count_syllables(word: str) -> int:
|
||||
"""
|
||||
Count syllables in a word using a heuristic approach.
|
||||
|
||||
Uses vowel group counting with adjustments for common patterns.
|
||||
|
||||
Args:
|
||||
word: The word to count syllables for.
|
||||
|
||||
Returns:
|
||||
Estimated syllable count (minimum 1 for non-empty words).
|
||||
"""
|
||||
if not word:
|
||||
return 0
|
||||
|
||||
@@ -25,33 +31,62 @@ def _count_syllables(word: str) -> int:
|
||||
if not word:
|
||||
return 0
|
||||
|
||||
# Count vowel groups
|
||||
vowel_groups = _VOWELS.findall(word)
|
||||
count = len(vowel_groups)
|
||||
|
||||
# Adjust for silent 'e' at end
|
||||
if word.endswith("e") and count > 1:
|
||||
count -= 1
|
||||
|
||||
# Adjust for 'le' ending (e.g., "table", "able")
|
||||
if word.endswith("le") and len(word) > 2 and word[-3] not in "aeiouy":
|
||||
count += 1
|
||||
|
||||
# Adjust for 'ed' ending when not adding syllable
|
||||
if word.endswith("ed") and len(word) > 2 and word[-3] not in "dt":
|
||||
count = max(count - 1, 1)
|
||||
|
||||
# Ensure at least 1 syllable for any word
|
||||
return max(count, 1)
|
||||
|
||||
|
||||
def _count_sentences(text: str) -> int:
|
||||
"""
|
||||
Count sentences in text.
|
||||
|
||||
Splits on sentence-ending punctuation (.!?).
|
||||
|
||||
Args:
|
||||
text: The text to count sentences in.
|
||||
|
||||
Returns:
|
||||
Number of sentences (minimum 1 for non-empty text).
|
||||
"""
|
||||
if not text or not text.strip():
|
||||
return 0
|
||||
|
||||
# Split on sentence endings and filter empty strings
|
||||
sentences = _SENTENCE_ENDINGS.split(text)
|
||||
# Filter out empty segments
|
||||
sentences = [s for s in sentences if s.strip()]
|
||||
|
||||
return max(len(sentences), 1)
|
||||
|
||||
|
||||
def _count_words(text: str) -> tuple[list[str], int]:
|
||||
"""
|
||||
Extract words from text and count them.
|
||||
|
||||
Args:
|
||||
text: The text to process.
|
||||
|
||||
Returns:
|
||||
Tuple of (word list, word count).
|
||||
"""
|
||||
# Extract words (sequences of letters and apostrophes)
|
||||
words = re.findall(r"[a-zA-Z']+", text)
|
||||
# Filter out standalone apostrophes
|
||||
words = [w for w in words if w.replace("'", "")]
|
||||
return words, len(words)
|
||||
|
||||
@@ -69,10 +104,12 @@ class Readability:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this metric."""
|
||||
return "readability"
|
||||
|
||||
@property
|
||||
def requires_reference(self) -> bool:
|
||||
"""Return whether this metric requires reference text."""
|
||||
return False
|
||||
|
||||
def score(
|
||||
@@ -90,31 +127,33 @@ class Readability:
|
||||
Returns:
|
||||
ReadabilityResult with Flesch-Kincaid scores.
|
||||
"""
|
||||
# Extract words and count
|
||||
words, word_count = _count_words(candidate)
|
||||
|
||||
# Handle empty or trivial text
|
||||
if word_count == 0:
|
||||
return ReadabilityResult(
|
||||
flesch_kincaid_grade=0.0,
|
||||
flesch_reading_ease=0.0,
|
||||
)
|
||||
|
||||
# Count sentences (ensure at least 1 to avoid division by zero)
|
||||
sentence_count = max(_count_sentences(candidate), 1)
|
||||
|
||||
# Count syllables
|
||||
syllable_count = sum(_count_syllables(word) for word in words)
|
||||
|
||||
# Compute ratios
|
||||
words_per_sentence = word_count / sentence_count
|
||||
syllables_per_word = syllable_count / word_count
|
||||
|
||||
grade_level = (
|
||||
FK_GRADE_WORDS_PER_SENTENCE * words_per_sentence
|
||||
+ FK_GRADE_SYLLABLES_PER_WORD * syllables_per_word
|
||||
- FK_GRADE_CONSTANT
|
||||
)
|
||||
# Flesch-Kincaid Grade Level
|
||||
# Formula: 0.39 * (words/sentences) + 11.8 * (syllables/words) - 15.59
|
||||
grade_level = 0.39 * words_per_sentence + 11.8 * syllables_per_word - 15.59
|
||||
|
||||
reading_ease = (
|
||||
FRE_CONSTANT
|
||||
- FRE_WORDS_PER_SENTENCE * words_per_sentence
|
||||
- FRE_SYLLABLES_PER_WORD * syllables_per_word
|
||||
)
|
||||
# Flesch Reading Ease
|
||||
# Formula: 206.835 - 1.015 * (words/sentences) - 84.6 * (syllables/words)
|
||||
reading_ease = 206.835 - 1.015 * words_per_sentence - 84.6 * syllables_per_word
|
||||
|
||||
return ReadabilityResult(
|
||||
flesch_kincaid_grade=grade_level,
|
||||
@@ -143,6 +182,7 @@ class Readability:
|
||||
for cand in candidates:
|
||||
results.append(self.score(cand))
|
||||
|
||||
# Compute aggregate statistics
|
||||
stats = {
|
||||
"flesch_kincaid_grade": AggregateStats.from_values(
|
||||
[r.flesch_kincaid_grade for r in results]
|
||||
|
||||
@@ -25,6 +25,7 @@ class BleuResult(BaseModel):
|
||||
|
||||
@property
|
||||
def score(self) -> float:
|
||||
"""Return the composite BLEU-4 score with brevity penalty."""
|
||||
return self.bleu4
|
||||
|
||||
|
||||
@@ -41,6 +42,7 @@ class LexicalResult(BaseModel):
|
||||
|
||||
@property
|
||||
def score(self) -> float:
|
||||
"""Return Jaccard similarity as the primary score."""
|
||||
return self.jaccard
|
||||
|
||||
|
||||
@@ -75,6 +77,7 @@ class RougeResult(BaseModel):
|
||||
|
||||
@property
|
||||
def score(self) -> float:
|
||||
"""Return ROUGE-L F-measure as the primary score."""
|
||||
return self.rouge_l.fmeasure
|
||||
|
||||
|
||||
@@ -91,6 +94,7 @@ class ReadabilityResult(BaseModel):
|
||||
|
||||
@property
|
||||
def score(self) -> float:
|
||||
"""Return Flesch reading ease as the primary score."""
|
||||
return self.flesch_reading_ease
|
||||
|
||||
|
||||
@@ -107,4 +111,5 @@ class SemanticResult(BaseModel):
|
||||
|
||||
@property
|
||||
def score(self) -> float:
|
||||
"""Return the primary score for this result."""
|
||||
return self.similarity
|
||||
|
||||
@@ -8,6 +8,7 @@ from veritext.metrics.results import RougeResult, RougeScore
|
||||
|
||||
|
||||
def _get_ngrams(tokens: list[str], n: int) -> Counter[tuple[str, ...]]:
|
||||
"""Extract n-grams from a list of tokens."""
|
||||
if n > len(tokens):
|
||||
return Counter()
|
||||
return Counter(tuple(tokens[i : i + n]) for i in range(len(tokens) - n + 1))
|
||||
@@ -17,6 +18,7 @@ def _ngram_overlap(
|
||||
candidate_ngrams: Counter[tuple[str, ...]],
|
||||
reference_ngrams: Counter[tuple[str, ...]],
|
||||
) -> int:
|
||||
"""Compute the overlap count between candidate and reference n-grams."""
|
||||
overlap = 0
|
||||
for ngram, count in candidate_ngrams.items():
|
||||
overlap += min(count, reference_ngrams.get(ngram, 0))
|
||||
@@ -70,11 +72,13 @@ def _lcs_length(seq1: list[str], seq2: list[str]) -> int:
|
||||
if not seq1 or not seq2:
|
||||
return 0
|
||||
|
||||
# Optimise by using shorter sequence for columns
|
||||
if len(seq1) < len(seq2):
|
||||
seq1, seq2 = seq2, seq1
|
||||
|
||||
m, n = len(seq1), len(seq2)
|
||||
|
||||
# Only need two rows at a time
|
||||
prev = [0] * (n + 1)
|
||||
curr = [0] * (n + 1)
|
||||
|
||||
@@ -120,6 +124,7 @@ def _compute_rouge_l(
|
||||
|
||||
|
||||
def _max_rouge_scores(scores: list[RougeScore]) -> RougeScore:
|
||||
"""Select the RougeScore with the highest F-measure from a list."""
|
||||
return max(scores, key=lambda s: s.fmeasure)
|
||||
|
||||
|
||||
@@ -142,10 +147,12 @@ class Rouge:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this metric."""
|
||||
return "rouge"
|
||||
|
||||
@property
|
||||
def requires_reference(self) -> bool:
|
||||
"""Return whether this metric requires reference text."""
|
||||
return True
|
||||
|
||||
def score(
|
||||
@@ -168,14 +175,18 @@ class Rouge:
|
||||
if reference is None:
|
||||
raise ValueError("ROUGE requires reference text")
|
||||
|
||||
# Normalise reference to list
|
||||
references = [reference] if isinstance(reference, str) else reference
|
||||
|
||||
# Tokenise
|
||||
candidate_tokens = self._tokeniser.tokenise(candidate)
|
||||
reference_token_lists = [self._tokeniser.tokenise(r) for r in references]
|
||||
|
||||
# Handle empty references
|
||||
if all(not ref for ref in reference_token_lists):
|
||||
raise ValueError("Reference text cannot be empty")
|
||||
|
||||
# Handle empty candidate
|
||||
if not candidate_tokens:
|
||||
return RougeResult(
|
||||
rouge1=RougeScore(precision=0.0, recall=0.0, fmeasure=0.0),
|
||||
@@ -183,6 +194,7 @@ class Rouge:
|
||||
rouge_l=RougeScore(precision=0.0, recall=0.0, fmeasure=0.0),
|
||||
)
|
||||
|
||||
# Compute scores for each reference and take max
|
||||
rouge1_scores = []
|
||||
rouge2_scores = []
|
||||
rouge_l_scores = []
|
||||
@@ -194,6 +206,7 @@ class Rouge:
|
||||
rouge2_scores.append(_compute_rouge_score(candidate_tokens, ref_tokens, 2))
|
||||
rouge_l_scores.append(_compute_rouge_l(candidate_tokens, ref_tokens))
|
||||
|
||||
# All references were empty after tokenisation
|
||||
if not rouge1_scores:
|
||||
raise ValueError("Reference text cannot be empty")
|
||||
|
||||
@@ -235,6 +248,7 @@ class Rouge:
|
||||
ref: str | list[str] = references[i]
|
||||
results.append(self.score(cand, ref))
|
||||
|
||||
# Compute aggregate statistics for each score type
|
||||
stats = {
|
||||
"rouge1_precision": AggregateStats.from_values(
|
||||
[r.rouge1.precision for r in results]
|
||||
|
||||
@@ -53,12 +53,14 @@ def validate_text(
|
||||
... max_reading_grade=8.0,
|
||||
... )
|
||||
"""
|
||||
# Validate that reference is provided for comparison metrics
|
||||
if any([min_bleu, min_rouge, min_semantic]) and reference is None:
|
||||
raise ValueError(
|
||||
"Reference text required for comparison metrics "
|
||||
"(min_bleu, min_rouge, min_semantic)"
|
||||
)
|
||||
|
||||
# Build list of validators from kwargs
|
||||
checks: list[Check] = []
|
||||
|
||||
if min_bleu is not None:
|
||||
@@ -72,6 +74,7 @@ def validate_text(
|
||||
checks.append(rouge(min_score=min_rouge))
|
||||
|
||||
if min_semantic is not None:
|
||||
# Lazy import to avoid loading sentence-transformers unless needed
|
||||
from veritext.validators import semantic
|
||||
|
||||
checks.append(semantic(min_score=min_semantic))
|
||||
@@ -99,6 +102,7 @@ def validate_text(
|
||||
if not checks:
|
||||
raise ValueError("At least one validation criterion must be specified")
|
||||
|
||||
# Run validation
|
||||
context = ValidationContext(reference=reference)
|
||||
validator = all_of(checks)
|
||||
result = validator.check(text, context)
|
||||
@@ -120,10 +124,12 @@ def _format_failure(text: str, result: ValidationResult) -> str:
|
||||
lines = ["Text validation failed:"]
|
||||
lines.append("")
|
||||
|
||||
# Show a preview of the text (truncated if long)
|
||||
preview = text[:100] + "..." if len(text) > 100 else text
|
||||
lines.append(f" Text: {preview!r}")
|
||||
lines.append("")
|
||||
|
||||
# List all failed checks with details
|
||||
lines.append(" Failed checks:")
|
||||
for check in result.failed_checks:
|
||||
lines.append(f" - {check.name}:")
|
||||
|
||||
@@ -7,6 +7,7 @@ from veritext.core.exceptions import DependencyError
|
||||
from veritext.metrics.base import AggregateStats, BatchResult
|
||||
from veritext.metrics.results import SemanticResult
|
||||
|
||||
# Default maximum cache size (number of embeddings to store)
|
||||
DEFAULT_CACHE_MAX_SIZE = 1000
|
||||
|
||||
|
||||
@@ -57,20 +58,33 @@ class SemanticSimilarity:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this metric."""
|
||||
return "semantic"
|
||||
|
||||
@property
|
||||
def requires_reference(self) -> bool:
|
||||
"""Return whether this metric requires reference text."""
|
||||
return True
|
||||
|
||||
def _get_embedding(self, text: str) -> Any:
|
||||
"""
|
||||
Get embedding for text, using LRU cache if available.
|
||||
|
||||
Args:
|
||||
text: The text to embed.
|
||||
|
||||
Returns:
|
||||
The embedding tensor.
|
||||
"""
|
||||
if self._cache is not None and text in self._cache:
|
||||
# Move to end to mark as recently used
|
||||
self._cache.move_to_end(text)
|
||||
return self._cache[text]
|
||||
|
||||
embedding = self._model.encode(text, convert_to_tensor=True)
|
||||
|
||||
if self._cache is not None:
|
||||
# Evict oldest entries if cache is full
|
||||
while len(self._cache) >= self._cache_max_size:
|
||||
self._cache.popitem(last=False)
|
||||
self._cache[text] = embedding
|
||||
@@ -78,9 +92,20 @@ class SemanticSimilarity:
|
||||
return embedding
|
||||
|
||||
def _cosine_similarity(self, embedding1: Any, embedding2: Any) -> float:
|
||||
"""
|
||||
Compute cosine similarity between two embeddings.
|
||||
|
||||
Args:
|
||||
embedding1: First embedding tensor.
|
||||
embedding2: Second embedding tensor.
|
||||
|
||||
Returns:
|
||||
Cosine similarity score (0.0 to 1.0).
|
||||
"""
|
||||
from sentence_transformers import util
|
||||
|
||||
similarity: float = util.cos_sim(embedding1, embedding2).item()
|
||||
# Clamp to [0, 1] as negative similarities are possible but not meaningful
|
||||
return max(0.0, min(1.0, similarity))
|
||||
|
||||
def score(
|
||||
@@ -105,21 +130,26 @@ class SemanticSimilarity:
|
||||
if reference is None:
|
||||
raise ValueError("Semantic similarity requires reference text")
|
||||
|
||||
# Normalise reference to list
|
||||
references = [reference] if isinstance(reference, str) else reference
|
||||
|
||||
if not references:
|
||||
raise ValueError("Reference text cannot be empty")
|
||||
|
||||
# Handle empty candidate
|
||||
candidate_stripped = candidate.strip()
|
||||
if not candidate_stripped:
|
||||
return SemanticResult(similarity=0.0, model=self._model_name)
|
||||
|
||||
# Handle empty references
|
||||
valid_references = [r for r in references if r.strip()]
|
||||
if not valid_references:
|
||||
raise ValueError("Reference text cannot be empty")
|
||||
|
||||
# Get candidate embedding
|
||||
candidate_embedding = self._get_embedding(candidate_stripped)
|
||||
|
||||
# Compute similarity against each reference, take maximum
|
||||
max_similarity = 0.0
|
||||
for ref in valid_references:
|
||||
ref_embedding = self._get_embedding(ref.strip())
|
||||
@@ -160,6 +190,7 @@ class SemanticSimilarity:
|
||||
ref: str | list[str] = references[i]
|
||||
results.append(self.score(cand, ref))
|
||||
|
||||
# Compute aggregate statistics
|
||||
stats = {
|
||||
"similarity": AggregateStats.from_values([r.similarity for r in results]),
|
||||
}
|
||||
@@ -167,5 +198,6 @@ class SemanticSimilarity:
|
||||
return BatchResult(results=results, count=len(results), stats=stats)
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
"""Clear the embedding cache."""
|
||||
if self._cache is not None:
|
||||
self._cache.clear()
|
||||
|
||||
@@ -35,6 +35,7 @@ from veritext.validators.metric import (
|
||||
)
|
||||
|
||||
|
||||
# Factory functions for clean API
|
||||
def bleu(
|
||||
min_score: float,
|
||||
variant: Literal[1, 2, 3, 4] = 4,
|
||||
|
||||
@@ -14,7 +14,9 @@ class Check(Protocol):
|
||||
"""
|
||||
|
||||
@property
|
||||
def name(self) -> str: ...
|
||||
def name(self) -> str:
|
||||
"""Return the name of this check."""
|
||||
...
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> CheckResult:
|
||||
"""Run the check and return a result.
|
||||
|
||||
@@ -33,6 +33,7 @@ class AllOf:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this composite check."""
|
||||
return "all_of"
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> ValidationResult:
|
||||
@@ -78,6 +79,7 @@ class AnyOf:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this composite check."""
|
||||
return "any_of"
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> ValidationResult:
|
||||
|
||||
@@ -61,6 +61,7 @@ class LengthValidator:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this check."""
|
||||
return "length"
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> CheckResult: # noqa: ARG002
|
||||
@@ -145,6 +146,7 @@ class ReadabilityValidator:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this check."""
|
||||
return "readability"
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> CheckResult: # noqa: ARG002
|
||||
@@ -211,16 +213,6 @@ class ReadabilityValidator:
|
||||
)
|
||||
|
||||
|
||||
def _compile_patterns(patterns: list[str], flags: int) -> list[re.Pattern[str]]:
|
||||
compiled = []
|
||||
for pattern in patterns:
|
||||
try:
|
||||
compiled.append(re.compile(pattern, flags))
|
||||
except re.error as e:
|
||||
raise InvalidThresholdError(f"Invalid regex pattern '{pattern}': {e}") from e
|
||||
return compiled
|
||||
|
||||
|
||||
class ContainsValidator:
|
||||
"""Validates text contains required patterns."""
|
||||
|
||||
@@ -245,10 +237,19 @@ class ContainsValidator:
|
||||
self._patterns = patterns
|
||||
self._case_sensitive = case_sensitive
|
||||
self._flags = 0 if case_sensitive else re.IGNORECASE
|
||||
self._compiled_patterns = _compile_patterns(patterns, self._flags)
|
||||
|
||||
self._compiled_patterns: list[re.Pattern[str]] = []
|
||||
for pattern in patterns:
|
||||
try:
|
||||
self._compiled_patterns.append(re.compile(pattern, self._flags))
|
||||
except re.error as e:
|
||||
raise InvalidThresholdError(
|
||||
f"Invalid regex pattern '{pattern}': {e}"
|
||||
) from e
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this check."""
|
||||
return "contains"
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> CheckResult: # noqa: ARG002
|
||||
@@ -309,10 +310,19 @@ class ExcludesValidator:
|
||||
self._patterns = patterns
|
||||
self._case_sensitive = case_sensitive
|
||||
self._flags = 0 if case_sensitive else re.IGNORECASE
|
||||
self._compiled_patterns = _compile_patterns(patterns, self._flags)
|
||||
|
||||
self._compiled_patterns: list[re.Pattern[str]] = []
|
||||
for pattern in patterns:
|
||||
try:
|
||||
self._compiled_patterns.append(re.compile(pattern, self._flags))
|
||||
except re.error as e:
|
||||
raise InvalidThresholdError(
|
||||
f"Invalid regex pattern '{pattern}': {e}"
|
||||
) from e
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this check."""
|
||||
return "excludes"
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> CheckResult: # noqa: ARG002
|
||||
|
||||
@@ -43,6 +43,7 @@ class BleuValidator:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this check."""
|
||||
return f"bleu-{self._variant}"
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> CheckResult:
|
||||
@@ -64,6 +65,7 @@ class BleuValidator:
|
||||
|
||||
result = self._metric.score(text, context.reference)
|
||||
|
||||
# Select the appropriate BLEU variant
|
||||
score_map = {
|
||||
1: result.bleu1,
|
||||
2: result.bleu2,
|
||||
@@ -128,6 +130,7 @@ class RougeValidator:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this check."""
|
||||
return f"rouge-{self._variant}"
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> CheckResult:
|
||||
@@ -149,6 +152,7 @@ class RougeValidator:
|
||||
|
||||
result = self._metric.score(text, context.reference)
|
||||
|
||||
# Select the appropriate ROUGE variant (use F-measure)
|
||||
score_map = {
|
||||
"1": result.rouge1.fmeasure,
|
||||
"2": result.rouge2.fmeasure,
|
||||
@@ -218,6 +222,7 @@ class LexicalValidator:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this check."""
|
||||
return "lexical"
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> CheckResult:
|
||||
@@ -239,6 +244,7 @@ class LexicalValidator:
|
||||
|
||||
result = self._metric.score(text, context.reference)
|
||||
|
||||
# Check each threshold that was specified
|
||||
failures = []
|
||||
if self._min_jaccard is not None and result.jaccard < self._min_jaccard:
|
||||
failures.append(
|
||||
@@ -265,6 +271,7 @@ class LexicalValidator:
|
||||
else:
|
||||
message = "Lexical similarity: " + "; ".join(failures)
|
||||
|
||||
# Build actual value dict
|
||||
actual = {"jaccard": result.jaccard, "token_overlap": result.token_overlap}
|
||||
threshold = {}
|
||||
if self._min_jaccard is not None:
|
||||
@@ -311,6 +318,7 @@ class SemanticValidator:
|
||||
)
|
||||
|
||||
self._min_score = min_score
|
||||
# Lazy import to avoid loading PyTorch unless needed
|
||||
from veritext.semantic.similarity import SemanticSimilarity
|
||||
|
||||
self._metric: SemanticSimilarity = SemanticSimilarity(
|
||||
@@ -319,6 +327,7 @@ class SemanticValidator:
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Return the name of this check."""
|
||||
return "semantic"
|
||||
|
||||
def check(self, text: str, context: ValidationContext) -> CheckResult:
|
||||
|
||||
@@ -7,14 +7,17 @@ from veritext.core.tokenisation import WordTokeniser
|
||||
|
||||
@pytest.fixture
|
||||
def word_tokeniser() -> WordTokeniser:
|
||||
"""Provide a default word tokeniser."""
|
||||
return WordTokeniser()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def word_tokeniser_no_lowercase() -> WordTokeniser:
|
||||
"""Provide a word tokeniser without lowercasing."""
|
||||
return WordTokeniser(lowercase=False)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def word_tokeniser_keep_punctuation() -> WordTokeniser:
|
||||
"""Provide a word tokeniser that keeps punctuation."""
|
||||
return WordTokeniser(remove_punctuation=False)
|
||||
|
||||
@@ -9,7 +9,10 @@ from veritext.benchmark.models import BenchmarkRun, RegressionReport
|
||||
|
||||
|
||||
class TestBenchmarkRun:
|
||||
"""Tests for BenchmarkRun model."""
|
||||
|
||||
def test_create_benchmark_run(self) -> None:
|
||||
"""BenchmarkRun can be created with required fields."""
|
||||
run = BenchmarkRun(
|
||||
id="test-id-123",
|
||||
benchmark_name="test-benchmark",
|
||||
@@ -27,6 +30,7 @@ class TestBenchmarkRun:
|
||||
assert run.metadata == {}
|
||||
|
||||
def test_create_with_metadata(self) -> None:
|
||||
"""BenchmarkRun can include optional metadata."""
|
||||
run = BenchmarkRun(
|
||||
id="test-id-456",
|
||||
benchmark_name="test-benchmark",
|
||||
@@ -40,6 +44,7 @@ class TestBenchmarkRun:
|
||||
assert run.metadata == {"git_sha": "abc123", "model_version": "gpt-4"}
|
||||
|
||||
def test_frozen_model(self) -> None:
|
||||
"""BenchmarkRun is immutable."""
|
||||
run = BenchmarkRun(
|
||||
id="test-id",
|
||||
benchmark_name="test",
|
||||
@@ -53,6 +58,7 @@ class TestBenchmarkRun:
|
||||
run.id = "new-id" # type: ignore[misc]
|
||||
|
||||
def test_serialisation(self) -> None:
|
||||
"""BenchmarkRun can be serialised to dict."""
|
||||
run = BenchmarkRun(
|
||||
id="test-id",
|
||||
benchmark_name="test",
|
||||
@@ -69,7 +75,10 @@ class TestBenchmarkRun:
|
||||
|
||||
|
||||
class TestRegressionReport:
|
||||
"""Tests for RegressionReport model."""
|
||||
|
||||
def test_no_regression_summary(self) -> None:
|
||||
"""Summary indicates no regression when detected is False."""
|
||||
report = RegressionReport(
|
||||
detected=False,
|
||||
baseline={"bleu4": 0.75, "rouge_l": 0.80},
|
||||
@@ -81,6 +90,7 @@ class TestRegressionReport:
|
||||
assert "No regression detected" in report.summary
|
||||
|
||||
def test_regression_summary(self) -> None:
|
||||
"""Summary lists regressed metrics when detected is True."""
|
||||
report = RegressionReport(
|
||||
detected=True,
|
||||
baseline={"bleu4": 0.75, "rouge_l": 0.80},
|
||||
@@ -94,7 +104,8 @@ class TestRegressionReport:
|
||||
assert "0.6500" in report.summary
|
||||
assert "baseline: 0.7500" in report.summary
|
||||
|
||||
def test_regression_within_tolerance(self) -> None:
|
||||
def test_regression_excludes_within_tolerance(self) -> None:
|
||||
"""Summary only shows metrics that exceed tolerance."""
|
||||
report = RegressionReport(
|
||||
detected=True,
|
||||
baseline={"bleu4": 0.75, "rouge_l": 0.80},
|
||||
@@ -109,6 +120,7 @@ class TestRegressionReport:
|
||||
assert "bleu4" in report.summary
|
||||
|
||||
def test_frozen_model(self) -> None:
|
||||
"""RegressionReport is immutable."""
|
||||
report = RegressionReport(
|
||||
detected=False,
|
||||
baseline={},
|
||||
@@ -121,6 +133,7 @@ class TestRegressionReport:
|
||||
report.detected = True # type: ignore[misc]
|
||||
|
||||
def test_tolerance_in_summary(self) -> None:
|
||||
"""Summary includes tolerance threshold."""
|
||||
report = RegressionReport(
|
||||
detected=True,
|
||||
baseline={"metric": 0.80},
|
||||
|
||||
@@ -13,6 +13,7 @@ def make_run(
|
||||
metrics: dict[str, float],
|
||||
day: int = 1,
|
||||
) -> BenchmarkRun:
|
||||
"""Helper to create a BenchmarkRun."""
|
||||
return BenchmarkRun(
|
||||
id=run_id,
|
||||
benchmark_name="test",
|
||||
@@ -24,11 +25,15 @@ def make_run(
|
||||
|
||||
|
||||
class TestComputeBaseline:
|
||||
"""Tests for baseline computation."""
|
||||
|
||||
def test_empty_runs(self) -> None:
|
||||
"""Returns empty baseline for empty runs list."""
|
||||
baseline = compute_baseline([])
|
||||
assert baseline == {}
|
||||
|
||||
def test_single_run(self) -> None:
|
||||
"""Single run produces baseline equal to that run's metrics."""
|
||||
runs = [make_run("r1", {"bleu4": 0.75, "rouge_l": 0.80})]
|
||||
|
||||
baseline = compute_baseline(runs)
|
||||
@@ -37,6 +42,7 @@ class TestComputeBaseline:
|
||||
assert baseline["rouge_l"] == 0.80
|
||||
|
||||
def test_multiple_runs_average(self) -> None:
|
||||
"""Baseline is the average of all runs in window."""
|
||||
runs = [
|
||||
make_run("r1", {"bleu4": 0.70}, day=3),
|
||||
make_run("r2", {"bleu4": 0.80}, day=2),
|
||||
@@ -48,6 +54,7 @@ class TestComputeBaseline:
|
||||
assert baseline["bleu4"] == pytest.approx(0.80) # (0.70+0.80+0.90)/3
|
||||
|
||||
def test_window_limits_runs(self) -> None:
|
||||
"""Only includes runs within the window size."""
|
||||
runs = [
|
||||
make_run("r1", {"bleu4": 0.70}, day=5), # most recent
|
||||
make_run("r2", {"bleu4": 0.80}, day=4),
|
||||
@@ -62,6 +69,7 @@ class TestComputeBaseline:
|
||||
assert baseline["bleu4"] == pytest.approx(0.80)
|
||||
|
||||
def test_partial_history(self) -> None:
|
||||
"""Works when fewer runs than window size exist."""
|
||||
runs = [
|
||||
make_run("r1", {"bleu4": 0.70}),
|
||||
make_run("r2", {"bleu4": 0.80}),
|
||||
@@ -73,6 +81,7 @@ class TestComputeBaseline:
|
||||
assert baseline["bleu4"] == pytest.approx(0.75)
|
||||
|
||||
def test_multiple_metrics(self) -> None:
|
||||
"""Computes baseline for all metrics present."""
|
||||
runs = [
|
||||
make_run("r1", {"bleu4": 0.70, "rouge_l": 0.75}),
|
||||
make_run("r2", {"bleu4": 0.80, "rouge_l": 0.85}),
|
||||
@@ -84,6 +93,7 @@ class TestComputeBaseline:
|
||||
assert baseline["rouge_l"] == pytest.approx(0.80)
|
||||
|
||||
def test_varying_metrics(self) -> None:
|
||||
"""Handles runs with different metric sets."""
|
||||
runs = [
|
||||
make_run("r1", {"bleu4": 0.70, "rouge_l": 0.75}),
|
||||
make_run("r2", {"bleu4": 0.80}), # No rouge_l
|
||||
@@ -98,7 +108,10 @@ class TestComputeBaseline:
|
||||
|
||||
|
||||
class TestDetectRegression:
|
||||
"""Tests for regression detection."""
|
||||
|
||||
def test_no_baseline(self) -> None:
|
||||
"""No regression when baseline is empty."""
|
||||
report = detect_regression(
|
||||
current={"bleu4": 0.70},
|
||||
baseline={},
|
||||
@@ -109,6 +122,7 @@ class TestDetectRegression:
|
||||
assert report.deltas == {}
|
||||
|
||||
def test_no_regression_stable(self) -> None:
|
||||
"""No regression when metrics are stable."""
|
||||
report = detect_regression(
|
||||
current={"bleu4": 0.75},
|
||||
baseline={"bleu4": 0.75},
|
||||
@@ -119,6 +133,7 @@ class TestDetectRegression:
|
||||
assert report.deltas["bleu4"] == pytest.approx(0.0)
|
||||
|
||||
def test_no_regression_improved(self) -> None:
|
||||
"""No regression when metrics improved."""
|
||||
report = detect_regression(
|
||||
current={"bleu4": 0.85},
|
||||
baseline={"bleu4": 0.75},
|
||||
@@ -129,6 +144,7 @@ class TestDetectRegression:
|
||||
assert report.deltas["bleu4"] == pytest.approx(0.10)
|
||||
|
||||
def test_no_regression_within_tolerance(self) -> None:
|
||||
"""No regression when drop is within tolerance."""
|
||||
report = detect_regression(
|
||||
current={"bleu4": 0.73},
|
||||
baseline={"bleu4": 0.75},
|
||||
@@ -139,6 +155,7 @@ class TestDetectRegression:
|
||||
assert report.deltas["bleu4"] == pytest.approx(-0.02)
|
||||
|
||||
def test_regression_detected(self) -> None:
|
||||
"""Regression detected when metric drops beyond tolerance."""
|
||||
report = detect_regression(
|
||||
current={"bleu4": 0.65},
|
||||
baseline={"bleu4": 0.75},
|
||||
@@ -149,6 +166,7 @@ class TestDetectRegression:
|
||||
assert report.deltas["bleu4"] == pytest.approx(-0.10)
|
||||
|
||||
def test_regression_at_tolerance_boundary(self) -> None:
|
||||
"""Drop at tolerance boundary is not a regression."""
|
||||
# Use a value clearly at the boundary (accounting for float precision)
|
||||
# The implementation checks delta < -tolerance (strictly less than)
|
||||
report = detect_regression(
|
||||
@@ -162,6 +180,7 @@ class TestDetectRegression:
|
||||
assert report.deltas["bleu4"] == 0.0
|
||||
|
||||
def test_regression_just_beyond_tolerance(self) -> None:
|
||||
"""Just beyond tolerance is a regression."""
|
||||
report = detect_regression(
|
||||
current={"bleu4": 0.6999},
|
||||
baseline={"bleu4": 0.75},
|
||||
@@ -172,6 +191,7 @@ class TestDetectRegression:
|
||||
assert report.detected
|
||||
|
||||
def test_multiple_metrics_any_regresses(self) -> None:
|
||||
"""Regression detected if any metric exceeds tolerance."""
|
||||
report = detect_regression(
|
||||
current={"bleu4": 0.65, "rouge_l": 0.80},
|
||||
baseline={"bleu4": 0.75, "rouge_l": 0.80},
|
||||
@@ -184,6 +204,7 @@ class TestDetectRegression:
|
||||
assert report.deltas["rouge_l"] == pytest.approx(0.0)
|
||||
|
||||
def test_report_contains_all_values(self) -> None:
|
||||
"""Report includes baseline, current, and deltas."""
|
||||
baseline = {"bleu4": 0.75, "rouge_l": 0.80}
|
||||
current = {"bleu4": 0.65, "rouge_l": 0.82}
|
||||
|
||||
@@ -196,6 +217,7 @@ class TestDetectRegression:
|
||||
assert "rouge_l" in report.deltas
|
||||
|
||||
def test_missing_metric_in_current(self) -> None:
|
||||
"""Missing metric in current treated as zero."""
|
||||
report = detect_regression(
|
||||
current={},
|
||||
baseline={"bleu4": 0.75},
|
||||
|
||||
@@ -11,11 +11,13 @@ from veritext.core.exceptions import RegressionDetectedError
|
||||
|
||||
@pytest.fixture
|
||||
def benchmark(tmp_path: Path) -> Benchmark:
|
||||
"""Create a Benchmark instance with temporary storage."""
|
||||
return Benchmark("test-suite", storage_path=tmp_path / "benchmarks")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_data() -> tuple[list[str], list[str]]:
|
||||
"""Sample candidates and references for testing."""
|
||||
candidates = [
|
||||
"The quick brown fox jumps over the lazy dog.",
|
||||
"A fast auburn fox leaps above the sleepy hound.",
|
||||
@@ -28,20 +30,27 @@ def sample_data() -> tuple[list[str], list[str]]:
|
||||
|
||||
|
||||
class TestBenchmarkInit:
|
||||
"""Tests for Benchmark initialisation."""
|
||||
|
||||
def test_creates_storage_directory(self, tmp_path: Path) -> None:
|
||||
"""Benchmark creates storage directory on init."""
|
||||
storage_path = tmp_path / "benchmarks"
|
||||
Benchmark("my-suite", storage_path=storage_path)
|
||||
|
||||
assert storage_path.exists()
|
||||
|
||||
def test_name_property(self, benchmark: Benchmark) -> None:
|
||||
"""Benchmark exposes its name."""
|
||||
assert benchmark.name == "test-suite"
|
||||
|
||||
|
||||
class TestEvaluate:
|
||||
"""Tests for the evaluate method."""
|
||||
|
||||
def test_evaluate_stores_run(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""Evaluate creates and stores a benchmark run."""
|
||||
candidates, references = sample_data
|
||||
|
||||
run = benchmark.evaluate(candidates, references)
|
||||
@@ -53,6 +62,7 @@ class TestEvaluate:
|
||||
def test_evaluate_returns_metrics(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""Evaluate computes default metrics."""
|
||||
candidates, references = sample_data
|
||||
|
||||
run = benchmark.evaluate(candidates, references)
|
||||
@@ -66,6 +76,7 @@ class TestEvaluate:
|
||||
def test_evaluate_custom_metrics(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""Evaluate can compute custom metrics."""
|
||||
candidates, references = sample_data
|
||||
|
||||
run = benchmark.evaluate(
|
||||
@@ -80,6 +91,7 @@ class TestEvaluate:
|
||||
def test_evaluate_with_metadata(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""Evaluate can include metadata."""
|
||||
candidates, references = sample_data
|
||||
|
||||
run = benchmark.evaluate(
|
||||
@@ -91,6 +103,7 @@ class TestEvaluate:
|
||||
def test_evaluate_stores_retrievable(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""Stored run can be retrieved."""
|
||||
candidates, references = sample_data
|
||||
run = benchmark.evaluate(candidates, references)
|
||||
|
||||
@@ -101,7 +114,10 @@ class TestEvaluate:
|
||||
|
||||
|
||||
class TestCheckRegression:
|
||||
"""Tests for regression checking."""
|
||||
|
||||
def test_check_no_runs(self, benchmark: Benchmark) -> None:
|
||||
"""No regression when no runs exist."""
|
||||
report = benchmark.check_regression()
|
||||
|
||||
assert not report.detected
|
||||
@@ -111,6 +127,7 @@ class TestCheckRegression:
|
||||
def test_check_single_run(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""No regression with single run (no baseline)."""
|
||||
candidates, references = sample_data
|
||||
benchmark.evaluate(candidates, references)
|
||||
|
||||
@@ -122,6 +139,7 @@ class TestCheckRegression:
|
||||
def test_check_stable_metrics(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""No regression when metrics are stable."""
|
||||
candidates, references = sample_data
|
||||
|
||||
# Run multiple times with same data
|
||||
@@ -132,6 +150,7 @@ class TestCheckRegression:
|
||||
assert not report.detected
|
||||
|
||||
def test_check_reports_regression(self, tmp_path: Path) -> None:
|
||||
"""Reports regression when metrics drop significantly."""
|
||||
benchmark = Benchmark("regress-test", storage_path=tmp_path / "benchmarks")
|
||||
|
||||
# First run with good metrics
|
||||
@@ -150,9 +169,12 @@ class TestCheckRegression:
|
||||
|
||||
|
||||
class TestAssertNoRegression:
|
||||
"""Tests for assert_no_regression method."""
|
||||
|
||||
def test_passes_when_stable(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""Does not raise when metrics are stable."""
|
||||
candidates, references = sample_data
|
||||
|
||||
for _ in range(3):
|
||||
@@ -162,6 +184,7 @@ class TestAssertNoRegression:
|
||||
benchmark.assert_no_regression()
|
||||
|
||||
def test_raises_on_regression(self, tmp_path: Path) -> None:
|
||||
"""Raises RegressionDetectedError when quality drops."""
|
||||
benchmark = Benchmark("regress-test", storage_path=tmp_path / "benchmarks")
|
||||
|
||||
# Establish baseline with perfect match
|
||||
@@ -177,13 +200,17 @@ class TestAssertNoRegression:
|
||||
|
||||
|
||||
class TestGetHistory:
|
||||
"""Tests for get_history method."""
|
||||
|
||||
def test_empty_history(self, benchmark: Benchmark) -> None:
|
||||
"""Returns empty list when no runs."""
|
||||
history = benchmark.get_history()
|
||||
assert history == []
|
||||
|
||||
def test_returns_runs(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""Returns benchmark runs."""
|
||||
candidates, references = sample_data
|
||||
|
||||
run1 = benchmark.evaluate(candidates, references)
|
||||
@@ -198,6 +225,7 @@ class TestGetHistory:
|
||||
def test_respects_limit(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""Respects limit parameter."""
|
||||
candidates, references = sample_data
|
||||
|
||||
for _ in range(5):
|
||||
@@ -209,6 +237,7 @@ class TestGetHistory:
|
||||
def test_default_limit(
|
||||
self, benchmark: Benchmark, sample_data: tuple[list[str], list[str]]
|
||||
) -> None:
|
||||
"""Default limit is 20."""
|
||||
candidates, references = sample_data
|
||||
|
||||
for _ in range(25):
|
||||
|
||||
@@ -14,16 +14,19 @@ from veritext.core.exceptions import StorageError
|
||||
|
||||
@pytest.fixture
|
||||
def db_path(tmp_path: Path) -> Path:
|
||||
"""Return a temporary database path."""
|
||||
return tmp_path / "benchmarks" / "test.db"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage(db_path: Path) -> BenchmarkStorage:
|
||||
"""Create a BenchmarkStorage instance."""
|
||||
return BenchmarkStorage(db_path)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_run() -> BenchmarkRun:
|
||||
"""Create a sample benchmark run."""
|
||||
return BenchmarkRun(
|
||||
id="run-001",
|
||||
benchmark_name="test-suite",
|
||||
@@ -36,17 +39,22 @@ def sample_run() -> BenchmarkRun:
|
||||
|
||||
|
||||
class TestDatabaseCreation:
|
||||
"""Tests for database initialisation."""
|
||||
|
||||
def test_creates_database_file(self, db_path: Path) -> None:
|
||||
"""Storage creates the database file on init."""
|
||||
assert not db_path.exists()
|
||||
BenchmarkStorage(db_path)
|
||||
assert db_path.exists()
|
||||
|
||||
def test_creates_parent_directories(self, tmp_path: Path) -> None:
|
||||
"""Storage creates parent directories if needed."""
|
||||
nested_path = tmp_path / "deep" / "nested" / "path" / "test.db"
|
||||
BenchmarkStorage(nested_path)
|
||||
assert nested_path.exists()
|
||||
|
||||
def test_creates_tables(self, db_path: Path) -> None:
|
||||
"""Storage creates required tables."""
|
||||
BenchmarkStorage(db_path)
|
||||
|
||||
conn = sqlite3.connect(str(db_path))
|
||||
@@ -58,6 +66,7 @@ class TestDatabaseCreation:
|
||||
assert "benchmark_metrics" in tables
|
||||
|
||||
def test_creates_index(self, db_path: Path) -> None:
|
||||
"""Storage creates index on benchmark_name and timestamp."""
|
||||
BenchmarkStorage(db_path)
|
||||
|
||||
conn = sqlite3.connect(str(db_path))
|
||||
@@ -69,9 +78,12 @@ class TestDatabaseCreation:
|
||||
|
||||
|
||||
class TestSaveRun:
|
||||
"""Tests for saving benchmark runs."""
|
||||
|
||||
def test_save_run(
|
||||
self, storage: BenchmarkStorage, sample_run: BenchmarkRun
|
||||
) -> None:
|
||||
"""Storage can save a benchmark run."""
|
||||
storage.save_run(sample_run)
|
||||
|
||||
runs = storage.get_runs("test-suite")
|
||||
@@ -81,6 +93,7 @@ class TestSaveRun:
|
||||
def test_save_preserves_all_fields(
|
||||
self, storage: BenchmarkStorage, sample_run: BenchmarkRun
|
||||
) -> None:
|
||||
"""Saved run preserves all fields correctly."""
|
||||
storage.save_run(sample_run)
|
||||
|
||||
runs = storage.get_runs("test-suite")
|
||||
@@ -97,12 +110,14 @@ class TestSaveRun:
|
||||
def test_save_duplicate_id_raises(
|
||||
self, storage: BenchmarkStorage, sample_run: BenchmarkRun
|
||||
) -> None:
|
||||
"""Saving a run with duplicate ID raises StorageError."""
|
||||
storage.save_run(sample_run)
|
||||
|
||||
with pytest.raises(StorageError, match="already exists"):
|
||||
storage.save_run(sample_run)
|
||||
|
||||
def test_save_run_empty_metadata(self, storage: BenchmarkStorage) -> None:
|
||||
"""Run with empty metadata saves correctly."""
|
||||
run = BenchmarkRun(
|
||||
id="run-no-meta",
|
||||
benchmark_name="test-suite",
|
||||
@@ -120,11 +135,15 @@ class TestSaveRun:
|
||||
|
||||
|
||||
class TestGetRuns:
|
||||
"""Tests for retrieving benchmark runs."""
|
||||
|
||||
def test_get_runs_empty_database(self, storage: BenchmarkStorage) -> None:
|
||||
"""Returns empty list for empty database."""
|
||||
runs = storage.get_runs("nonexistent")
|
||||
assert runs == []
|
||||
|
||||
def test_get_runs_filters_by_name(self, storage: BenchmarkStorage) -> None:
|
||||
"""Returns only runs matching the benchmark name."""
|
||||
run1 = BenchmarkRun(
|
||||
id="run-1",
|
||||
benchmark_name="suite-a",
|
||||
@@ -154,6 +173,7 @@ class TestGetRuns:
|
||||
assert runs_b[0].id == "run-2"
|
||||
|
||||
def test_get_runs_ordered_by_timestamp(self, storage: BenchmarkStorage) -> None:
|
||||
"""Returns runs ordered by timestamp, most recent first."""
|
||||
run_old = BenchmarkRun(
|
||||
id="run-old",
|
||||
benchmark_name="test",
|
||||
@@ -180,6 +200,7 @@ class TestGetRuns:
|
||||
assert runs[1].id == "run-old"
|
||||
|
||||
def test_get_runs_with_limit(self, storage: BenchmarkStorage) -> None:
|
||||
"""Respects limit parameter."""
|
||||
for i in range(5):
|
||||
run = BenchmarkRun(
|
||||
id=f"run-{i}",
|
||||
@@ -196,11 +217,15 @@ class TestGetRuns:
|
||||
|
||||
|
||||
class TestGetLatestRun:
|
||||
"""Tests for getting the latest run."""
|
||||
|
||||
def test_get_latest_run_empty(self, storage: BenchmarkStorage) -> None:
|
||||
"""Returns None for empty database."""
|
||||
result = storage.get_latest_run("nonexistent")
|
||||
assert result is None
|
||||
|
||||
def test_get_latest_run(self, storage: BenchmarkStorage) -> None:
|
||||
"""Returns the most recent run."""
|
||||
run_old = BenchmarkRun(
|
||||
id="run-old",
|
||||
benchmark_name="test",
|
||||
@@ -227,7 +252,10 @@ class TestGetLatestRun:
|
||||
|
||||
|
||||
class TestConcurrentAccess:
|
||||
"""Tests for concurrent database access."""
|
||||
|
||||
def test_concurrent_writes(self, db_path: Path) -> None:
|
||||
"""Multiple threads can write concurrently with WAL mode."""
|
||||
errors: list[Exception] = []
|
||||
|
||||
def write_run(run_id: int) -> None:
|
||||
@@ -258,6 +286,7 @@ class TestConcurrentAccess:
|
||||
assert len(runs) == 10
|
||||
|
||||
def test_wal_mode_enabled(self, db_path: Path) -> None:
|
||||
"""Database uses WAL journal mode."""
|
||||
BenchmarkStorage(db_path)
|
||||
|
||||
conn = sqlite3.connect(str(db_path))
|
||||
|
||||
@@ -10,7 +10,10 @@ runner = CliRunner()
|
||||
|
||||
|
||||
class TestBenchmarkRun:
|
||||
"""Tests for benchmark run command."""
|
||||
|
||||
def test_benchmark_run_basic(self, tmp_path: Path) -> None:
|
||||
"""Test basic benchmark run."""
|
||||
data_file = tmp_path / "data.jsonl"
|
||||
data_file.write_text(
|
||||
'{"candidate": "hello world today", "reference": "hello world today"}\n'
|
||||
@@ -39,6 +42,7 @@ class TestBenchmarkRun:
|
||||
assert "bleu4:" in result.stdout
|
||||
|
||||
def test_benchmark_run_file_not_found(self, tmp_path: Path) -> None:
|
||||
"""Test benchmark run with non-existent file."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
[
|
||||
@@ -55,6 +59,7 @@ class TestBenchmarkRun:
|
||||
assert "Error" in result.stdout
|
||||
|
||||
def test_benchmark_run_creates_storage(self, tmp_path: Path) -> None:
|
||||
"""Test that benchmark run creates storage directory."""
|
||||
data_file = tmp_path / "data.jsonl"
|
||||
data_file.write_text('{"candidate": "hello", "reference": "hello"}')
|
||||
storage_path = tmp_path / "new_benchmarks"
|
||||
@@ -76,7 +81,10 @@ class TestBenchmarkRun:
|
||||
|
||||
|
||||
class TestBenchmarkShow:
|
||||
"""Tests for benchmark show command."""
|
||||
|
||||
def test_benchmark_show_no_runs(self, tmp_path: Path) -> None:
|
||||
"""Test showing benchmark with no runs."""
|
||||
storage_path = tmp_path / "benchmarks"
|
||||
storage_path.mkdir()
|
||||
|
||||
@@ -94,6 +102,7 @@ class TestBenchmarkShow:
|
||||
assert "No benchmark runs found" in result.stdout
|
||||
|
||||
def test_benchmark_show_with_runs(self, tmp_path: Path) -> None:
|
||||
"""Test showing benchmark history with runs."""
|
||||
# First create some runs
|
||||
data_file = tmp_path / "data.jsonl"
|
||||
data_file.write_text('{"candidate": "hello world", "reference": "hello world"}')
|
||||
@@ -129,6 +138,7 @@ class TestBenchmarkShow:
|
||||
assert "Benchmark History" in result.stdout
|
||||
|
||||
def test_benchmark_show_limit(self, tmp_path: Path) -> None:
|
||||
"""Test showing limited benchmark history."""
|
||||
data_file = tmp_path / "data.jsonl"
|
||||
data_file.write_text('{"candidate": "hello", "reference": "hello"}')
|
||||
storage_path = tmp_path / "benchmarks"
|
||||
@@ -165,7 +175,10 @@ class TestBenchmarkShow:
|
||||
|
||||
|
||||
class TestBenchmarkCheck:
|
||||
"""Tests for benchmark check command."""
|
||||
|
||||
def test_benchmark_check_no_regression(self, tmp_path: Path) -> None:
|
||||
"""Test checking for regression with no regression."""
|
||||
data_file = tmp_path / "data.jsonl"
|
||||
data_file.write_text(
|
||||
'{"candidate": "hello world today", "reference": "hello world today"}'
|
||||
@@ -202,6 +215,7 @@ class TestBenchmarkCheck:
|
||||
assert "No regression detected" in result.stdout
|
||||
|
||||
def test_benchmark_check_with_regression(self, tmp_path: Path) -> None:
|
||||
"""Test checking for regression when regression occurs."""
|
||||
storage_path = tmp_path / "benchmarks"
|
||||
|
||||
# First run with good data
|
||||
@@ -257,6 +271,7 @@ class TestBenchmarkCheck:
|
||||
assert "Regression detected" in result.stdout
|
||||
|
||||
def test_benchmark_check_custom_tolerance(self, tmp_path: Path) -> None:
|
||||
"""Test checking regression with custom tolerance."""
|
||||
data_file = tmp_path / "data.jsonl"
|
||||
data_file.write_text('{"candidate": "hello", "reference": "hello"}')
|
||||
storage_path = tmp_path / "benchmarks"
|
||||
@@ -291,7 +306,10 @@ class TestBenchmarkCheck:
|
||||
|
||||
|
||||
class TestBenchmarkHelp:
|
||||
"""Tests for benchmark help output."""
|
||||
|
||||
def test_benchmark_help(self) -> None:
|
||||
"""Test benchmark help output."""
|
||||
result = runner.invoke(app, ["benchmark", "--help"])
|
||||
assert result.exit_code == 0
|
||||
assert "run" in result.stdout
|
||||
@@ -299,17 +317,20 @@ class TestBenchmarkHelp:
|
||||
assert "check" in result.stdout
|
||||
|
||||
def test_benchmark_run_help(self) -> None:
|
||||
"""Test benchmark run help output."""
|
||||
result = runner.invoke(app, ["benchmark", "run", "--help"])
|
||||
assert result.exit_code == 0
|
||||
assert "--file" in result.stdout
|
||||
assert "--metrics" in result.stdout
|
||||
|
||||
def test_benchmark_show_help(self) -> None:
|
||||
"""Test benchmark show help output."""
|
||||
result = runner.invoke(app, ["benchmark", "show", "--help"])
|
||||
assert result.exit_code == 0
|
||||
assert "--last" in result.stdout
|
||||
|
||||
def test_benchmark_check_help(self) -> None:
|
||||
"""Test benchmark check help output."""
|
||||
result = runner.invoke(app, ["benchmark", "check", "--help"])
|
||||
assert result.exit_code == 0
|
||||
assert "--tolerance" in result.stdout
|
||||
|
||||
@@ -13,22 +13,28 @@ from veritext.cli.formatters import (
|
||||
|
||||
|
||||
class TestFormatValidationTable:
|
||||
"""Tests for format_validation_table function."""
|
||||
|
||||
def test_format_empty_results(self) -> None:
|
||||
"""Test formatting empty results."""
|
||||
table = format_validation_table({})
|
||||
assert table.title == "Validation Results"
|
||||
assert table.row_count == 0
|
||||
|
||||
def test_format_single_metric(self) -> None:
|
||||
"""Test formatting a single metric."""
|
||||
results = {"bleu4": 0.8523}
|
||||
table = format_validation_table(results)
|
||||
assert table.row_count == 1
|
||||
|
||||
def test_format_multiple_metrics(self) -> None:
|
||||
"""Test formatting multiple metrics."""
|
||||
results = {"bleu4": 0.85, "rouge_l": 0.92, "jaccard": 0.75}
|
||||
table = format_validation_table(results)
|
||||
assert table.row_count == 3
|
||||
|
||||
def test_format_with_threshold(self) -> None:
|
||||
"""Test formatting with threshold for pass/fail."""
|
||||
results = {"bleu4": 0.85, "rouge_l": 0.45}
|
||||
table = format_validation_table(results, threshold=0.5)
|
||||
# Should have 3 columns: Metric, Score, Status
|
||||
@@ -36,11 +42,15 @@ class TestFormatValidationTable:
|
||||
|
||||
|
||||
class TestFormatValidationJson:
|
||||
"""Tests for format_validation_json function."""
|
||||
|
||||
def test_format_empty_results(self) -> None:
|
||||
"""Test formatting empty results as JSON."""
|
||||
result = format_validation_json({})
|
||||
assert result == "{}"
|
||||
|
||||
def test_format_results(self) -> None:
|
||||
"""Test formatting results as JSON."""
|
||||
results = {"bleu4": 0.85, "rouge_l": 0.92}
|
||||
result = format_validation_json(results)
|
||||
assert '"bleu4": 0.85' in result
|
||||
@@ -48,11 +58,15 @@ class TestFormatValidationJson:
|
||||
|
||||
|
||||
class TestFormatValidationSimple:
|
||||
"""Tests for format_validation_simple function."""
|
||||
|
||||
def test_format_empty_results(self) -> None:
|
||||
"""Test formatting empty results as simple text."""
|
||||
result = format_validation_simple({})
|
||||
assert result == ""
|
||||
|
||||
def test_format_results(self) -> None:
|
||||
"""Test formatting results as simple text."""
|
||||
results = {"bleu4": 0.8523, "rouge_l": 0.9234}
|
||||
result = format_validation_simple(results)
|
||||
assert "bleu4: 0.8523" in result
|
||||
@@ -60,11 +74,15 @@ class TestFormatValidationSimple:
|
||||
|
||||
|
||||
class TestFormatBenchmarkHistory:
|
||||
"""Tests for format_benchmark_history function."""
|
||||
|
||||
def test_format_empty_history(self) -> None:
|
||||
"""Test formatting empty benchmark history."""
|
||||
table = format_benchmark_history([])
|
||||
assert table.title == "Benchmark History"
|
||||
|
||||
def test_format_single_run(self) -> None:
|
||||
"""Test formatting a single benchmark run."""
|
||||
run = BenchmarkRun(
|
||||
id="test-id",
|
||||
benchmark_name="test",
|
||||
@@ -77,6 +95,7 @@ class TestFormatBenchmarkHistory:
|
||||
assert table.row_count == 1
|
||||
|
||||
def test_format_multiple_runs(self) -> None:
|
||||
"""Test formatting multiple benchmark runs."""
|
||||
runs = [
|
||||
BenchmarkRun(
|
||||
id=f"test-id-{i}",
|
||||
@@ -93,7 +112,10 @@ class TestFormatBenchmarkHistory:
|
||||
|
||||
|
||||
class TestFormatRegressionReport:
|
||||
"""Tests for format_regression_report function."""
|
||||
|
||||
def test_format_no_regression(self) -> None:
|
||||
"""Test formatting report with no regression."""
|
||||
report = RegressionReport(
|
||||
detected=False,
|
||||
baseline={"rouge_l": 0.85},
|
||||
@@ -106,6 +128,7 @@ class TestFormatRegressionReport:
|
||||
assert panel.border_style == "green"
|
||||
|
||||
def test_format_with_regression(self) -> None:
|
||||
"""Test formatting report with regression detected."""
|
||||
report = RegressionReport(
|
||||
detected=True,
|
||||
baseline={"rouge_l": 0.85, "bleu4": 0.72},
|
||||
|
||||
@@ -9,14 +9,20 @@ from veritext.cli.readers import TextPair, read_jsonl, read_paired_jsonl
|
||||
|
||||
|
||||
class TestTextPair:
|
||||
"""Tests for TextPair dataclass."""
|
||||
|
||||
def test_create_text_pair(self) -> None:
|
||||
"""Test creating a TextPair."""
|
||||
pair = TextPair(candidate="hello", reference="world")
|
||||
assert pair.candidate == "hello"
|
||||
assert pair.reference == "world"
|
||||
|
||||
|
||||
class TestReadJsonl:
|
||||
"""Tests for read_jsonl function."""
|
||||
|
||||
def test_read_valid_jsonl(self, tmp_path: Path) -> None:
|
||||
"""Test reading a valid JSONL file."""
|
||||
data = [
|
||||
{"candidate": "foo", "reference": "bar"},
|
||||
{"candidate": "baz", "reference": "qux"},
|
||||
@@ -33,6 +39,7 @@ class TestReadJsonl:
|
||||
assert pairs[1].reference == "qux"
|
||||
|
||||
def test_read_empty_file(self, tmp_path: Path) -> None:
|
||||
"""Test reading an empty JSONL file."""
|
||||
jsonl_file = tmp_path / "empty.jsonl"
|
||||
jsonl_file.write_text("")
|
||||
|
||||
@@ -41,6 +48,7 @@ class TestReadJsonl:
|
||||
assert pairs == []
|
||||
|
||||
def test_read_file_with_blank_lines(self, tmp_path: Path) -> None:
|
||||
"""Test reading a JSONL file with blank lines."""
|
||||
jsonl_file = tmp_path / "data.jsonl"
|
||||
content = '{"candidate": "a", "reference": "b"}\n\n{"candidate": "c", "reference": "d"}\n'
|
||||
jsonl_file.write_text(content)
|
||||
@@ -50,10 +58,12 @@ class TestReadJsonl:
|
||||
assert len(pairs) == 2
|
||||
|
||||
def test_read_file_not_found(self, tmp_path: Path) -> None:
|
||||
"""Test reading a non-existent file."""
|
||||
with pytest.raises(FileNotFoundError):
|
||||
read_jsonl(tmp_path / "nonexistent.jsonl")
|
||||
|
||||
def test_read_invalid_json(self, tmp_path: Path) -> None:
|
||||
"""Test reading a file with invalid JSON."""
|
||||
jsonl_file = tmp_path / "invalid.jsonl"
|
||||
jsonl_file.write_text("not valid json")
|
||||
|
||||
@@ -61,6 +71,7 @@ class TestReadJsonl:
|
||||
read_jsonl(jsonl_file)
|
||||
|
||||
def test_read_missing_candidate_key(self, tmp_path: Path) -> None:
|
||||
"""Test reading a file missing the candidate key."""
|
||||
jsonl_file = tmp_path / "data.jsonl"
|
||||
jsonl_file.write_text('{"reference": "bar"}')
|
||||
|
||||
@@ -68,6 +79,7 @@ class TestReadJsonl:
|
||||
read_jsonl(jsonl_file)
|
||||
|
||||
def test_read_missing_reference_key(self, tmp_path: Path) -> None:
|
||||
"""Test reading a file missing the reference key."""
|
||||
jsonl_file = tmp_path / "data.jsonl"
|
||||
jsonl_file.write_text('{"candidate": "foo"}')
|
||||
|
||||
@@ -76,7 +88,10 @@ class TestReadJsonl:
|
||||
|
||||
|
||||
class TestReadPairedJsonl:
|
||||
"""Tests for read_paired_jsonl function."""
|
||||
|
||||
def test_read_paired_valid(self, tmp_path: Path) -> None:
|
||||
"""Test reading valid paired JSONL files."""
|
||||
candidates_file = tmp_path / "candidates.jsonl"
|
||||
references_file = tmp_path / "references.jsonl"
|
||||
|
||||
@@ -92,6 +107,7 @@ class TestReadPairedJsonl:
|
||||
assert pairs[1].reference == "qux"
|
||||
|
||||
def test_read_paired_length_mismatch(self, tmp_path: Path) -> None:
|
||||
"""Test reading paired files with different lengths."""
|
||||
candidates_file = tmp_path / "candidates.jsonl"
|
||||
references_file = tmp_path / "references.jsonl"
|
||||
|
||||
@@ -102,6 +118,7 @@ class TestReadPairedJsonl:
|
||||
read_paired_jsonl(candidates_file, references_file)
|
||||
|
||||
def test_read_paired_candidates_not_found(self, tmp_path: Path) -> None:
|
||||
"""Test reading when candidates file doesn't exist."""
|
||||
references_file = tmp_path / "references.jsonl"
|
||||
references_file.write_text('{"text": "baz"}')
|
||||
|
||||
@@ -109,6 +126,7 @@ class TestReadPairedJsonl:
|
||||
read_paired_jsonl(tmp_path / "nonexistent.jsonl", references_file)
|
||||
|
||||
def test_read_paired_references_not_found(self, tmp_path: Path) -> None:
|
||||
"""Test reading when references file doesn't exist."""
|
||||
candidates_file = tmp_path / "candidates.jsonl"
|
||||
candidates_file.write_text('{"text": "foo"}')
|
||||
|
||||
@@ -116,6 +134,7 @@ class TestReadPairedJsonl:
|
||||
read_paired_jsonl(candidates_file, tmp_path / "nonexistent.jsonl")
|
||||
|
||||
def test_read_paired_missing_text_key(self, tmp_path: Path) -> None:
|
||||
"""Test reading paired files with missing text key."""
|
||||
candidates_file = tmp_path / "candidates.jsonl"
|
||||
references_file = tmp_path / "references.jsonl"
|
||||
|
||||
|
||||
@@ -11,7 +11,10 @@ runner = CliRunner()
|
||||
|
||||
|
||||
class TestValidateInline:
|
||||
"""Tests for inline validation mode."""
|
||||
|
||||
def test_validate_inline_basic(self) -> None:
|
||||
"""Test basic inline validation."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
[
|
||||
@@ -27,6 +30,7 @@ class TestValidateInline:
|
||||
assert "bleu4" in result.stdout
|
||||
|
||||
def test_validate_inline_with_rouge(self) -> None:
|
||||
"""Test inline validation with ROUGE metric."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
[
|
||||
@@ -42,6 +46,7 @@ class TestValidateInline:
|
||||
assert "rouge_l" in result.stdout
|
||||
|
||||
def test_validate_inline_with_lexical(self) -> None:
|
||||
"""Test inline validation with lexical metric."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
[
|
||||
@@ -58,6 +63,7 @@ class TestValidateInline:
|
||||
assert "token_overlap" in result.stdout
|
||||
|
||||
def test_validate_inline_json_output(self) -> None:
|
||||
"""Test inline validation with JSON output."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
[
|
||||
@@ -76,6 +82,7 @@ class TestValidateInline:
|
||||
assert "bleu4" in data
|
||||
|
||||
def test_validate_inline_simple_output(self) -> None:
|
||||
"""Test inline validation with simple output."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
[
|
||||
@@ -93,6 +100,7 @@ class TestValidateInline:
|
||||
assert "rouge_l:" in result.stdout
|
||||
|
||||
def test_validate_inline_missing_reference(self) -> None:
|
||||
"""Test inline validation without reference."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
["validate", "hello world", "-m", "bleu"],
|
||||
@@ -101,6 +109,7 @@ class TestValidateInline:
|
||||
assert "Error" in result.stdout
|
||||
|
||||
def test_validate_inline_invalid_metric(self) -> None:
|
||||
"""Test inline validation with invalid metric."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
["validate", "hello", "-r", "world", "-m", "invalid_metric"],
|
||||
@@ -110,7 +119,10 @@ class TestValidateInline:
|
||||
|
||||
|
||||
class TestValidateFile:
|
||||
"""Tests for file-based validation mode."""
|
||||
|
||||
def test_validate_file_basic(self, tmp_path: Path) -> None:
|
||||
"""Test basic file-based validation."""
|
||||
data_file = tmp_path / "data.jsonl"
|
||||
data_file.write_text(
|
||||
'{"candidate": "hello world today", "reference": "hello world today"}\n'
|
||||
@@ -126,6 +138,7 @@ class TestValidateFile:
|
||||
assert "Evaluated 2 text pairs" in result.stdout
|
||||
|
||||
def test_validate_file_not_found(self) -> None:
|
||||
"""Test file-based validation with non-existent file."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
["validate", "-f", "/nonexistent/file.jsonl", "-m", "bleu"],
|
||||
@@ -134,6 +147,7 @@ class TestValidateFile:
|
||||
assert "Error" in result.stdout
|
||||
|
||||
def test_validate_paired_files(self, tmp_path: Path) -> None:
|
||||
"""Test validation with separate candidate and reference files."""
|
||||
candidates_file = tmp_path / "candidates.jsonl"
|
||||
references_file = tmp_path / "references.jsonl"
|
||||
|
||||
@@ -161,7 +175,10 @@ class TestValidateFile:
|
||||
|
||||
|
||||
class TestValidateOptions:
|
||||
"""Tests for validate command options."""
|
||||
|
||||
def test_validate_with_threshold(self) -> None:
|
||||
"""Test validation with threshold option."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
[
|
||||
@@ -180,6 +197,7 @@ class TestValidateOptions:
|
||||
assert "Status" in result.stdout or "PASS" in result.stdout
|
||||
|
||||
def test_validate_invalid_output_format(self) -> None:
|
||||
"""Test validation with invalid output format."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
[
|
||||
@@ -197,6 +215,7 @@ class TestValidateOptions:
|
||||
assert "Invalid output format" in result.stdout
|
||||
|
||||
def test_validate_multiple_metrics(self) -> None:
|
||||
"""Test validation with multiple metrics."""
|
||||
result = runner.invoke(
|
||||
app,
|
||||
[
|
||||
|
||||
@@ -8,51 +8,66 @@ from veritext.core.config import VeritextSettings, get_settings
|
||||
|
||||
|
||||
class TestVeritextSettings:
|
||||
"""Tests for VeritextSettings."""
|
||||
|
||||
def test_default_log_level(self) -> None:
|
||||
"""Test default log level is INFO."""
|
||||
settings = VeritextSettings()
|
||||
assert settings.log_level == "INFO"
|
||||
|
||||
def test_default_log_format(self) -> None:
|
||||
"""Test default log format is console."""
|
||||
settings = VeritextSettings()
|
||||
assert settings.log_format == "console"
|
||||
|
||||
def test_default_benchmark_path(self) -> None:
|
||||
"""Test default benchmark storage path."""
|
||||
settings = VeritextSettings()
|
||||
assert settings.benchmark_storage_path == Path("benchmarks")
|
||||
|
||||
def test_default_tokeniser_lowercase(self) -> None:
|
||||
"""Test default tokeniser lowercase setting."""
|
||||
settings = VeritextSettings()
|
||||
assert settings.tokeniser_lowercase is True
|
||||
|
||||
def test_default_remove_punctuation(self) -> None:
|
||||
def test_default_tokeniser_remove_punctuation(self) -> None:
|
||||
"""Test default tokeniser remove punctuation setting."""
|
||||
settings = VeritextSettings()
|
||||
assert settings.tokeniser_remove_punctuation is True
|
||||
|
||||
def test_default_semantic_model(self) -> None:
|
||||
"""Test default semantic model name."""
|
||||
settings = VeritextSettings()
|
||||
assert settings.semantic_model == "all-MiniLM-L6-v2"
|
||||
|
||||
def test_default_semantic_cache_enabled(self) -> None:
|
||||
"""Test semantic cache is enabled by default."""
|
||||
settings = VeritextSettings()
|
||||
assert settings.semantic_cache_embeddings is True
|
||||
|
||||
def test_env_var_override(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Test environment variable overrides default settings."""
|
||||
monkeypatch.setenv("VERITEXT_LOG_LEVEL", "DEBUG")
|
||||
settings = VeritextSettings()
|
||||
assert settings.log_level == "DEBUG"
|
||||
|
||||
def test_env_var_override_log_format(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Test environment variable overrides log format."""
|
||||
monkeypatch.setenv("VERITEXT_LOG_FORMAT", "json")
|
||||
settings = VeritextSettings()
|
||||
assert settings.log_format == "json"
|
||||
|
||||
|
||||
class TestGetSettings:
|
||||
"""Tests for get_settings function."""
|
||||
|
||||
def test_get_settings_returns_instance(self) -> None:
|
||||
"""Test get_settings returns a VeritextSettings instance."""
|
||||
settings = get_settings()
|
||||
assert isinstance(settings, VeritextSettings)
|
||||
|
||||
def test_get_settings_returns_valid_defaults(self) -> None:
|
||||
"""Test get_settings returns instance with valid defaults."""
|
||||
settings = get_settings()
|
||||
assert settings.log_level in ("DEBUG", "INFO", "WARNING", "ERROR")
|
||||
assert settings.log_format in ("console", "json")
|
||||
|
||||
@@ -4,11 +4,15 @@ from veritext.core.logging import configure_logging, get_logger
|
||||
|
||||
|
||||
class TestGetLogger:
|
||||
"""Tests for get_logger function."""
|
||||
|
||||
def test_get_logger_returns_logger(self) -> None:
|
||||
"""Test get_logger returns a logger instance."""
|
||||
logger = get_logger()
|
||||
assert logger is not None
|
||||
|
||||
def test_get_logger_default_name(self) -> None:
|
||||
"""Test get_logger uses 'veritext' as default name."""
|
||||
logger = get_logger()
|
||||
# The logger should be a bound logger from structlog
|
||||
assert hasattr(logger, "info")
|
||||
@@ -17,28 +21,35 @@ class TestGetLogger:
|
||||
assert hasattr(logger, "error")
|
||||
|
||||
def test_get_logger_custom_name(self) -> None:
|
||||
"""Test get_logger respects custom name parameter."""
|
||||
logger = get_logger("custom.module")
|
||||
assert logger is not None
|
||||
assert hasattr(logger, "info")
|
||||
|
||||
|
||||
class TestConfigureLogging:
|
||||
"""Tests for configure_logging function."""
|
||||
|
||||
def test_configure_logging_console_format(self) -> None:
|
||||
"""Test configure_logging with console format does not raise."""
|
||||
configure_logging(level="INFO", log_format="console")
|
||||
logger = get_logger()
|
||||
assert logger is not None
|
||||
|
||||
def test_configure_logging_json_format(self) -> None:
|
||||
"""Test configure_logging with json format does not raise."""
|
||||
configure_logging(level="DEBUG", log_format="json")
|
||||
logger = get_logger()
|
||||
assert logger is not None
|
||||
|
||||
def test_configure_logging_uses_defaults(self) -> None:
|
||||
"""Test configure_logging uses settings defaults when not provided."""
|
||||
configure_logging()
|
||||
logger = get_logger()
|
||||
assert logger is not None
|
||||
|
||||
def test_configure_logging_different_levels(self) -> None:
|
||||
"""Test configure_logging accepts different log levels."""
|
||||
for level in ("DEBUG", "INFO", "WARNING", "ERROR"):
|
||||
configure_logging(level=level)
|
||||
logger = get_logger()
|
||||
|
||||
@@ -9,41 +9,52 @@ if TYPE_CHECKING:
|
||||
|
||||
|
||||
class TestWordTokeniser:
|
||||
"""Tests for WordTokeniser."""
|
||||
|
||||
def test_basic_tokenisation(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test basic word tokenisation."""
|
||||
tokens = word_tokeniser.tokenise("The cat sat on the mat")
|
||||
assert tokens == ["the", "cat", "sat", "on", "the", "mat"]
|
||||
|
||||
def test_lowercasing(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test that tokens are lowercased by default."""
|
||||
tokens = word_tokeniser.tokenise("Hello WORLD")
|
||||
assert tokens == ["hello", "world"]
|
||||
|
||||
def test_no_lowercasing(self, word_tokeniser_no_lowercase: WordTokeniser) -> None:
|
||||
"""Test tokenisation without lowercasing."""
|
||||
tokens = word_tokeniser_no_lowercase.tokenise("Hello WORLD")
|
||||
assert tokens == ["Hello", "WORLD"]
|
||||
|
||||
def test_punctuation_removal(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test that punctuation is removed by default."""
|
||||
tokens = word_tokeniser.tokenise("Hello, world! How are you?")
|
||||
assert tokens == ["hello", "world", "how", "are", "you"]
|
||||
|
||||
def test_keep_punctuation(
|
||||
self, word_tokeniser_keep_punctuation: WordTokeniser
|
||||
) -> None:
|
||||
"""Test tokenisation keeping punctuation."""
|
||||
tokens = word_tokeniser_keep_punctuation.tokenise("Hello, world!")
|
||||
assert tokens == ["hello,", "world!"]
|
||||
|
||||
def test_empty_string(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test that empty string returns empty list."""
|
||||
tokens = word_tokeniser.tokenise("")
|
||||
assert tokens == []
|
||||
|
||||
def test_whitespace_only(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test that whitespace-only string returns empty list."""
|
||||
tokens = word_tokeniser.tokenise(" \t\n ")
|
||||
assert tokens == []
|
||||
|
||||
def test_multiple_spaces(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test handling of multiple spaces between words."""
|
||||
tokens = word_tokeniser.tokenise("hello world")
|
||||
assert tokens == ["hello", "world"]
|
||||
|
||||
def test_unicode_nfc_normalisation(self) -> None:
|
||||
"""Test NFC Unicode normalisation."""
|
||||
# 'é' can be composed (U+00E9) or decomposed (e + U+0301)
|
||||
composed = "caf\u00e9" # café with composed é
|
||||
decomposed = "cafe\u0301" # café with decomposed é
|
||||
@@ -57,33 +68,40 @@ class TestWordTokeniser:
|
||||
assert tokens_composed == ["café"]
|
||||
|
||||
def test_unicode_emoji(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test handling of emoji characters."""
|
||||
# Emoji are removed as punctuation by default
|
||||
tokens = word_tokeniser.tokenise("Hello 👋 world 🌍")
|
||||
assert tokens == ["hello", "world"]
|
||||
|
||||
def test_unicode_non_latin(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test handling of non-Latin scripts."""
|
||||
tokens = word_tokeniser.tokenise("日本語 テスト")
|
||||
assert tokens == ["日本語", "テスト"]
|
||||
|
||||
def test_unicode_mixed_scripts(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test handling of mixed scripts."""
|
||||
tokens = word_tokeniser.tokenise("Hello 世界 Bonjour мир")
|
||||
assert tokens == ["hello", "世界", "bonjour", "мир"]
|
||||
|
||||
def test_numbers_preserved(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test that numbers are preserved."""
|
||||
tokens = word_tokeniser.tokenise("I have 42 apples")
|
||||
assert tokens == ["i", "have", "42", "apples"]
|
||||
|
||||
def test_contractions(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test handling of contractions (apostrophe replaced with space)."""
|
||||
tokens = word_tokeniser.tokenise("I can't don't won't")
|
||||
# Apostrophe is replaced with space, splitting the words
|
||||
assert tokens == ["i", "can", "t", "don", "t", "won", "t"]
|
||||
|
||||
def test_hyphenated_words(self, word_tokeniser: WordTokeniser) -> None:
|
||||
"""Test handling of hyphenated words."""
|
||||
tokens = word_tokeniser.tokenise("state-of-the-art")
|
||||
# Hyphens are removed as punctuation
|
||||
assert tokens == ["state", "of", "the", "art"]
|
||||
|
||||
def test_custom_normalisation_form(self) -> None:
|
||||
"""Test custom Unicode normalisation form."""
|
||||
tokeniser = WordTokeniser(normalisation_form="NFKC")
|
||||
# NFKC normalises compatibility characters
|
||||
# ™ (U+2122) is compatibility equivalent to 'TM'
|
||||
@@ -92,7 +110,10 @@ class TestWordTokeniser:
|
||||
|
||||
|
||||
class TestTokeniserProtocol:
|
||||
"""Tests for Tokeniser protocol compliance."""
|
||||
|
||||
def test_word_tokeniser_implements_protocol(self) -> None:
|
||||
"""Test that WordTokeniser implements the Tokeniser protocol."""
|
||||
tokeniser: Tokeniser = WordTokeniser()
|
||||
# Check that it has the required method
|
||||
assert hasattr(tokeniser, "tokenise")
|
||||
|
||||
@@ -7,34 +7,44 @@ from veritext.core.types import CheckResult, ValidationContext, ValidationResult
|
||||
|
||||
|
||||
class TestValidationContext:
|
||||
"""Tests for ValidationContext."""
|
||||
|
||||
def test_create_empty_context(self) -> None:
|
||||
"""Test creating a context with no reference."""
|
||||
context = ValidationContext()
|
||||
assert context.reference is None
|
||||
assert context.metadata == {}
|
||||
|
||||
def test_create_with_single_reference(self) -> None:
|
||||
"""Test creating a context with a single reference string."""
|
||||
context = ValidationContext(reference="The cat sat on the mat.")
|
||||
assert context.reference == "The cat sat on the mat."
|
||||
|
||||
def test_create_with_multiple_references(self) -> None:
|
||||
"""Test creating a context with multiple reference strings."""
|
||||
references = ["The cat sat on the mat.", "A cat was sitting on a mat."]
|
||||
context = ValidationContext(reference=references)
|
||||
assert context.reference == references
|
||||
assert len(context.reference) == 2
|
||||
|
||||
def test_create_with_metadata(self) -> None:
|
||||
"""Test creating a context with metadata."""
|
||||
metadata = {"source": "test", "timestamp": "2024-01-01"}
|
||||
context = ValidationContext(metadata=metadata)
|
||||
assert context.metadata == metadata
|
||||
|
||||
def test_context_is_immutable(self) -> None:
|
||||
"""Test that ValidationContext is immutable (frozen)."""
|
||||
context = ValidationContext(reference="test")
|
||||
with pytest.raises(PydanticValidationError):
|
||||
context.reference = "new value" # type: ignore[misc]
|
||||
|
||||
|
||||
class TestCheckResult:
|
||||
"""Tests for CheckResult."""
|
||||
|
||||
def test_create_passing_result(self) -> None:
|
||||
"""Test creating a passing check result."""
|
||||
result = CheckResult(
|
||||
name="bleu",
|
||||
passed=True,
|
||||
@@ -49,6 +59,7 @@ class TestCheckResult:
|
||||
assert "0.85" in result.message
|
||||
|
||||
def test_create_failing_result(self) -> None:
|
||||
"""Test creating a failing check result."""
|
||||
result = CheckResult(
|
||||
name="length",
|
||||
passed=False,
|
||||
@@ -62,6 +73,7 @@ class TestCheckResult:
|
||||
assert result.threshold == 500
|
||||
|
||||
def test_create_result_without_threshold(self) -> None:
|
||||
"""Test creating a result for checks without thresholds."""
|
||||
result = CheckResult(
|
||||
name="contains",
|
||||
passed=True,
|
||||
@@ -72,6 +84,7 @@ class TestCheckResult:
|
||||
assert result.threshold is None
|
||||
|
||||
def test_result_is_immutable(self) -> None:
|
||||
"""Test that CheckResult is immutable (frozen)."""
|
||||
result = CheckResult(
|
||||
name="test",
|
||||
passed=True,
|
||||
@@ -83,7 +96,10 @@ class TestCheckResult:
|
||||
|
||||
|
||||
class TestValidationResult:
|
||||
"""Tests for ValidationResult."""
|
||||
|
||||
def test_all_checks_passed(self) -> None:
|
||||
"""Test validation result when all checks pass."""
|
||||
checks = [
|
||||
CheckResult(name="bleu", passed=True, actual=0.8, message="OK"),
|
||||
CheckResult(name="length", passed=True, actual=100, message="OK"),
|
||||
@@ -96,6 +112,7 @@ class TestValidationResult:
|
||||
assert "All checks passed" in result.failure_summary
|
||||
|
||||
def test_some_checks_failed(self) -> None:
|
||||
"""Test validation result when some checks fail."""
|
||||
checks = [
|
||||
CheckResult(name="bleu", passed=True, actual=0.8, message="OK"),
|
||||
CheckResult(
|
||||
@@ -122,6 +139,7 @@ class TestValidationResult:
|
||||
assert result.failed_checks[1].name == "readability"
|
||||
|
||||
def test_failure_summary_format(self) -> None:
|
||||
"""Test the format of failure summary."""
|
||||
checks = [
|
||||
CheckResult(
|
||||
name="bleu",
|
||||
@@ -140,6 +158,7 @@ class TestValidationResult:
|
||||
assert "BLEU score 0.5 < threshold 0.7" in summary
|
||||
|
||||
def test_empty_checks_list(self) -> None:
|
||||
"""Test validation result with no checks."""
|
||||
result = ValidationResult(passed=True, checks=[])
|
||||
assert result.passed is True
|
||||
assert result.checks == []
|
||||
@@ -147,11 +166,13 @@ class TestValidationResult:
|
||||
assert "All checks passed" in result.failure_summary
|
||||
|
||||
def test_result_is_immutable(self) -> None:
|
||||
"""Test that ValidationResult is immutable (frozen)."""
|
||||
result = ValidationResult(passed=True, checks=[])
|
||||
with pytest.raises(PydanticValidationError):
|
||||
result.passed = False # type: ignore[misc]
|
||||
|
||||
def test_check_actual_can_be_any_type(self) -> None:
|
||||
"""Test that CheckResult.actual can hold any type."""
|
||||
# List
|
||||
result1 = CheckResult(
|
||||
name="contains",
|
||||
|
||||
@@ -6,17 +6,23 @@ from veritext.metrics import Bleu, BleuResult
|
||||
|
||||
|
||||
class TestBleu:
|
||||
"""Tests for the Bleu metric class."""
|
||||
|
||||
@pytest.fixture
|
||||
def bleu(self) -> Bleu:
|
||||
"""Provide a BLEU metric instance."""
|
||||
return Bleu()
|
||||
|
||||
def test_name(self, bleu: Bleu) -> None:
|
||||
"""Test that name returns 'bleu'."""
|
||||
assert bleu.name == "bleu"
|
||||
|
||||
def test_requires_reference(self, bleu: Bleu) -> None:
|
||||
"""Test that BLEU requires reference text."""
|
||||
assert bleu.requires_reference is True
|
||||
|
||||
def test_identical_texts(self, bleu: Bleu) -> None:
|
||||
"""Test that identical texts produce perfect scores."""
|
||||
text = "The cat sat on the mat"
|
||||
result = bleu.score(text, text)
|
||||
|
||||
@@ -27,6 +33,7 @@ class TestBleu:
|
||||
assert result.brevity_penalty == 1.0
|
||||
|
||||
def test_no_overlap(self, bleu: Bleu) -> None:
|
||||
"""Test that texts with no overlap produce zero scores."""
|
||||
candidate = "The cat sat"
|
||||
reference = "A dog runs fast"
|
||||
result = bleu.score(candidate, reference)
|
||||
@@ -37,34 +44,44 @@ class TestBleu:
|
||||
assert result.bleu4 == 0.0
|
||||
|
||||
def test_partial_overlap(self, bleu: Bleu) -> None:
|
||||
"""Test partial overlap produces intermediate scores."""
|
||||
candidate = "The cat sat on the mat"
|
||||
reference = "The cat is on the floor"
|
||||
result = bleu.score(candidate, reference)
|
||||
|
||||
# Should have some overlap but not perfect
|
||||
assert 0.0 < result.bleu1 < 1.0
|
||||
assert 0.0 < result.bleu2 < 1.0
|
||||
|
||||
def test_brevity_penalty_applied(self, bleu: Bleu) -> None:
|
||||
"""Test that brevity penalty is applied for short candidates."""
|
||||
candidate = "cat"
|
||||
reference = "The cat sat on the mat"
|
||||
result = bleu.score(candidate, reference)
|
||||
|
||||
# Candidate is much shorter, so brevity penalty should be < 1
|
||||
assert result.brevity_penalty < 1.0
|
||||
|
||||
def test_no_brevity_penalty_long_cand(self, bleu: Bleu) -> None:
|
||||
def test_no_brevity_penalty_for_long_candidate(self, bleu: Bleu) -> None:
|
||||
"""Test that no brevity penalty for longer candidates."""
|
||||
candidate = "The cat sat on the mat today"
|
||||
reference = "The cat sat"
|
||||
result = bleu.score(candidate, reference)
|
||||
|
||||
# Candidate is longer, so no brevity penalty
|
||||
assert result.brevity_penalty == 1.0
|
||||
|
||||
def test_single_word_texts(self, bleu: Bleu) -> None:
|
||||
"""Test BLEU with single word texts."""
|
||||
result = bleu.score("cat", "cat")
|
||||
|
||||
assert result.bleu1 == 1.0
|
||||
# Higher order n-grams can't be computed for single words
|
||||
# but with geometric mean, they still produce a score
|
||||
assert result.brevity_penalty == 1.0
|
||||
|
||||
def test_empty_candidate(self, bleu: Bleu) -> None:
|
||||
"""Test that empty candidate returns zero scores."""
|
||||
result = bleu.score("", "The cat sat")
|
||||
|
||||
assert result.bleu1 == 0.0
|
||||
@@ -74,20 +91,24 @@ class TestBleu:
|
||||
assert result.brevity_penalty == 0.0
|
||||
|
||||
def test_whitespace_only_candidate(self, bleu: Bleu) -> None:
|
||||
"""Test that whitespace-only candidate returns zero scores."""
|
||||
result = bleu.score(" \t\n ", "The cat sat")
|
||||
|
||||
assert result.bleu1 == 0.0
|
||||
assert result.bleu4 == 0.0
|
||||
|
||||
def test_empty_reference_raises(self, bleu: Bleu) -> None:
|
||||
"""Test that empty reference raises ValueError."""
|
||||
with pytest.raises(ValueError, match="cannot be empty"):
|
||||
bleu.score("The cat sat", "")
|
||||
|
||||
def test_none_reference_raises(self, bleu: Bleu) -> None:
|
||||
"""Test that None reference raises ValueError."""
|
||||
with pytest.raises(ValueError, match="requires reference"):
|
||||
bleu.score("The cat sat", None)
|
||||
|
||||
def test_multiple_references(self, bleu: Bleu) -> None:
|
||||
"""Test BLEU with multiple references uses max n-gram matches."""
|
||||
candidate = "The cat sat on the mat"
|
||||
references = [
|
||||
"The cat is on the floor",
|
||||
@@ -95,9 +116,11 @@ class TestBleu:
|
||||
]
|
||||
result = bleu.score(candidate, references)
|
||||
|
||||
# Should get perfect score due to exact match reference
|
||||
assert result.bleu4 == 1.0
|
||||
|
||||
def test_multiple_references_partial(self, bleu: Bleu) -> None:
|
||||
"""Test multiple references with partial matches."""
|
||||
candidate = "The quick brown fox"
|
||||
references = [
|
||||
"The fast brown fox",
|
||||
@@ -105,27 +128,35 @@ class TestBleu:
|
||||
]
|
||||
result = bleu.score(candidate, references)
|
||||
|
||||
# Should benefit from both references
|
||||
assert result.bleu1 > 0.0
|
||||
|
||||
def test_result_score_property(self, bleu: Bleu) -> None:
|
||||
"""Test that result.score returns bleu4."""
|
||||
result = bleu.score("The cat sat", "The cat sat")
|
||||
assert result.score == result.bleu4
|
||||
|
||||
def test_case_insensitivity(self, bleu: Bleu) -> None:
|
||||
"""Test that BLEU is case insensitive by default."""
|
||||
result = bleu.score("THE CAT SAT ON THE MAT", "the cat sat on the mat")
|
||||
assert result.bleu4 == 1.0
|
||||
|
||||
def test_punctuation_ignored(self, bleu: Bleu) -> None:
|
||||
"""Test that punctuation is ignored by default."""
|
||||
result = bleu.score("The cat sat on the mat.", "The cat sat on the mat!")
|
||||
assert result.bleu4 == 1.0
|
||||
|
||||
|
||||
class TestBleuBatch:
|
||||
"""Tests for BLEU batch scoring."""
|
||||
|
||||
@pytest.fixture
|
||||
def bleu(self) -> Bleu:
|
||||
"""Provide a BLEU metric instance."""
|
||||
return Bleu()
|
||||
|
||||
def test_batch_score_basic(self, bleu: Bleu) -> None:
|
||||
"""Test basic batch scoring."""
|
||||
candidates = ["The cat sat on the mat", "A quick brown dog runs fast"]
|
||||
references = ["The cat sat on the mat", "A quick brown dog runs fast"]
|
||||
result = bleu.batch_score(candidates, references)
|
||||
@@ -135,6 +166,7 @@ class TestBleuBatch:
|
||||
assert all(r.bleu4 == 1.0 for r in result.results)
|
||||
|
||||
def test_batch_score_statistics(self, bleu: Bleu) -> None:
|
||||
"""Test that batch scoring computes statistics."""
|
||||
candidates = ["The cat sat", "A different text entirely"]
|
||||
references = ["The cat sat", "The cat sat"]
|
||||
result = bleu.batch_score(candidates, references)
|
||||
@@ -149,6 +181,7 @@ class TestBleuBatch:
|
||||
assert stats.min <= stats.mean <= stats.max
|
||||
|
||||
def test_batch_score_percentiles(self, bleu: Bleu) -> None:
|
||||
"""Test that batch scoring computes percentiles."""
|
||||
candidates = ["a", "b", "c", "d", "e"]
|
||||
references = ["a", "b", "c", "d", "e"]
|
||||
result = bleu.batch_score(candidates, references)
|
||||
@@ -160,14 +193,17 @@ class TestBleuBatch:
|
||||
assert 95 in stats.percentiles
|
||||
|
||||
def test_batch_score_none_references_raises(self, bleu: Bleu) -> None:
|
||||
"""Test that batch scoring raises for None references."""
|
||||
with pytest.raises(ValueError, match="requires reference"):
|
||||
bleu.batch_score(["text"], None)
|
||||
|
||||
def test_batch_score_length_mismatch_raises(self, bleu: Bleu) -> None:
|
||||
"""Test that batch scoring raises for mismatched lengths."""
|
||||
with pytest.raises(ValueError, match="must match"):
|
||||
bleu.batch_score(["a", "b"], ["a"])
|
||||
|
||||
def test_batch_score_multi_refs(self, bleu: Bleu) -> None:
|
||||
def test_batch_score_with_multiple_references(self, bleu: Bleu) -> None:
|
||||
"""Test batch scoring with multiple references per candidate."""
|
||||
candidates = [
|
||||
"The cat sat on the mat",
|
||||
"A quick brown dog runs fast",
|
||||
@@ -184,7 +220,10 @@ class TestBleuBatch:
|
||||
|
||||
|
||||
class TestBleuResult:
|
||||
"""Tests for BleuResult type."""
|
||||
|
||||
def test_frozen(self) -> None:
|
||||
"""Test that BleuResult is frozen."""
|
||||
from pydantic import ValidationError
|
||||
|
||||
result = BleuResult(
|
||||
@@ -194,6 +233,7 @@ class TestBleuResult:
|
||||
result.bleu1 = 0.6 # type: ignore[misc]
|
||||
|
||||
def test_score_property(self) -> None:
|
||||
"""Test that score property returns bleu4."""
|
||||
result = BleuResult(
|
||||
bleu1=0.9, bleu2=0.8, bleu3=0.7, bleu4=0.6, brevity_penalty=1.0
|
||||
)
|
||||
|
||||
@@ -6,17 +6,23 @@ from veritext.metrics import Lexical, LexicalResult
|
||||
|
||||
|
||||
class TestLexical:
|
||||
"""Tests for the Lexical metric class."""
|
||||
|
||||
@pytest.fixture
|
||||
def lexical(self) -> Lexical:
|
||||
"""Provide a lexical metric instance."""
|
||||
return Lexical()
|
||||
|
||||
def test_name(self, lexical: Lexical) -> None:
|
||||
"""Test that name returns 'lexical'."""
|
||||
assert lexical.name == "lexical"
|
||||
|
||||
def test_requires_reference(self, lexical: Lexical) -> None:
|
||||
"""Test that lexical requires reference text."""
|
||||
assert lexical.requires_reference is True
|
||||
|
||||
def test_identical_texts(self, lexical: Lexical) -> None:
|
||||
"""Test that identical texts produce perfect scores."""
|
||||
text = "The cat sat on the mat"
|
||||
result = lexical.score(text, text)
|
||||
|
||||
@@ -24,6 +30,7 @@ class TestLexical:
|
||||
assert result.token_overlap == 1.0
|
||||
|
||||
def test_no_overlap(self, lexical: Lexical) -> None:
|
||||
"""Test that texts with no overlap produce zero scores."""
|
||||
candidate = "apple banana cherry"
|
||||
reference = "dog elephant fox"
|
||||
result = lexical.score(candidate, reference)
|
||||
@@ -32,6 +39,7 @@ class TestLexical:
|
||||
assert result.token_overlap == 0.0
|
||||
|
||||
def test_partial_overlap_jaccard(self, lexical: Lexical) -> None:
|
||||
"""Test Jaccard with partial overlap."""
|
||||
candidate = "the cat sat"
|
||||
reference = "the dog sat"
|
||||
result = lexical.score(candidate, reference)
|
||||
@@ -41,6 +49,7 @@ class TestLexical:
|
||||
assert result.jaccard == 0.5
|
||||
|
||||
def test_partial_overlap_token_overlap(self, lexical: Lexical) -> None:
|
||||
"""Test token overlap with partial overlap."""
|
||||
candidate = "the cat sat"
|
||||
reference = "the dog sat"
|
||||
result = lexical.score(candidate, reference)
|
||||
@@ -51,6 +60,7 @@ class TestLexical:
|
||||
assert abs(result.token_overlap - 2 / 3) < 1e-10
|
||||
|
||||
def test_candidate_subset_of_reference(self, lexical: Lexical) -> None:
|
||||
"""Test when candidate is a subset of reference."""
|
||||
candidate = "the cat"
|
||||
reference = "the cat sat on the mat"
|
||||
result = lexical.score(candidate, reference)
|
||||
@@ -61,6 +71,7 @@ class TestLexical:
|
||||
assert result.jaccard < 1.0
|
||||
|
||||
def test_reference_subset_of_candidate(self, lexical: Lexical) -> None:
|
||||
"""Test when reference is a subset of candidate."""
|
||||
candidate = "the cat sat on the mat"
|
||||
reference = "the cat"
|
||||
result = lexical.score(candidate, reference)
|
||||
@@ -71,26 +82,31 @@ class TestLexical:
|
||||
assert result.token_overlap < 1.0
|
||||
|
||||
def test_empty_candidate(self, lexical: Lexical) -> None:
|
||||
"""Test that empty candidate returns zero scores."""
|
||||
result = lexical.score("", "The cat sat")
|
||||
|
||||
assert result.jaccard == 0.0
|
||||
assert result.token_overlap == 0.0
|
||||
|
||||
def test_whitespace_only_candidate(self, lexical: Lexical) -> None:
|
||||
"""Test that whitespace-only candidate returns zero scores."""
|
||||
result = lexical.score(" \t\n ", "The cat sat")
|
||||
|
||||
assert result.jaccard == 0.0
|
||||
assert result.token_overlap == 0.0
|
||||
|
||||
def test_empty_reference_raises(self, lexical: Lexical) -> None:
|
||||
"""Test that empty reference raises ValueError."""
|
||||
with pytest.raises(ValueError, match="cannot be empty"):
|
||||
lexical.score("The cat sat", "")
|
||||
|
||||
def test_none_reference_raises(self, lexical: Lexical) -> None:
|
||||
"""Test that None reference raises ValueError."""
|
||||
with pytest.raises(ValueError, match="requires reference"):
|
||||
lexical.score("The cat sat", None)
|
||||
|
||||
def test_multiple_references_uses_first(self, lexical: Lexical) -> None:
|
||||
"""Test that multiple references uses the first one."""
|
||||
candidate = "the cat sat"
|
||||
references = ["the dog ran", "the cat sat"] # First differs
|
||||
result = lexical.score(candidate, references)
|
||||
@@ -99,16 +115,19 @@ class TestLexical:
|
||||
assert result.jaccard < 1.0
|
||||
|
||||
def test_case_insensitivity(self, lexical: Lexical) -> None:
|
||||
"""Test that lexical is case insensitive by default."""
|
||||
result = lexical.score("THE CAT SAT", "the cat sat")
|
||||
assert result.jaccard == 1.0
|
||||
assert result.token_overlap == 1.0
|
||||
|
||||
def test_punctuation_ignored(self, lexical: Lexical) -> None:
|
||||
"""Test that punctuation is ignored by default."""
|
||||
result = lexical.score("The cat sat.", "The cat sat!")
|
||||
assert result.jaccard == 1.0
|
||||
assert result.token_overlap == 1.0
|
||||
|
||||
def test_repeated_tokens(self, lexical: Lexical) -> None:
|
||||
"""Test handling of repeated tokens."""
|
||||
candidate = "the the the"
|
||||
reference = "the cat"
|
||||
result = lexical.score(candidate, reference)
|
||||
@@ -121,11 +140,15 @@ class TestLexical:
|
||||
|
||||
|
||||
class TestLexicalBatch:
|
||||
"""Tests for lexical batch scoring."""
|
||||
|
||||
@pytest.fixture
|
||||
def lexical(self) -> Lexical:
|
||||
"""Provide a lexical metric instance."""
|
||||
return Lexical()
|
||||
|
||||
def test_batch_score_basic(self, lexical: Lexical) -> None:
|
||||
"""Test basic batch scoring."""
|
||||
candidates = ["The cat sat", "A dog runs"]
|
||||
references = ["The cat sat", "A dog runs"]
|
||||
result = lexical.batch_score(candidates, references)
|
||||
@@ -135,6 +158,7 @@ class TestLexicalBatch:
|
||||
assert all(r.jaccard == 1.0 for r in result.results)
|
||||
|
||||
def test_batch_score_statistics(self, lexical: Lexical) -> None:
|
||||
"""Test that batch scoring computes statistics."""
|
||||
candidates = ["The cat sat", "Completely different words"]
|
||||
references = ["The cat sat", "The cat sat"]
|
||||
result = lexical.batch_score(candidates, references)
|
||||
@@ -151,6 +175,7 @@ class TestLexicalBatch:
|
||||
assert result.stats["jaccard"].mean == 0.5
|
||||
|
||||
def test_batch_score_percentiles(self, lexical: Lexical) -> None:
|
||||
"""Test that batch scoring computes percentiles."""
|
||||
candidates = ["a", "b", "c", "d", "e"]
|
||||
references = ["a", "b", "c", "d", "e"]
|
||||
result = lexical.batch_score(candidates, references)
|
||||
@@ -162,16 +187,21 @@ class TestLexicalBatch:
|
||||
assert 95 in stats.percentiles
|
||||
|
||||
def test_batch_score_none_references_raises(self, lexical: Lexical) -> None:
|
||||
"""Test that batch scoring raises for None references."""
|
||||
with pytest.raises(ValueError, match="requires reference"):
|
||||
lexical.batch_score(["text"], None)
|
||||
|
||||
def test_batch_score_length_mismatch_raises(self, lexical: Lexical) -> None:
|
||||
"""Test that batch scoring raises for mismatched lengths."""
|
||||
with pytest.raises(ValueError, match="must match"):
|
||||
lexical.batch_score(["a", "b"], ["a"])
|
||||
|
||||
|
||||
class TestLexicalResult:
|
||||
"""Tests for LexicalResult type."""
|
||||
|
||||
def test_frozen(self) -> None:
|
||||
"""Test that LexicalResult is frozen."""
|
||||
from pydantic import ValidationError
|
||||
|
||||
result = LexicalResult(jaccard=0.5, token_overlap=0.7)
|
||||
@@ -179,6 +209,7 @@ class TestLexicalResult:
|
||||
result.jaccard = 0.6 # type: ignore[misc]
|
||||
|
||||
def test_values(self) -> None:
|
||||
"""Test that values are stored correctly."""
|
||||
result = LexicalResult(jaccard=0.5, token_overlap=0.7)
|
||||
assert result.jaccard == 0.5
|
||||
assert result.token_overlap == 0.7
|
||||
|
||||
@@ -6,17 +6,23 @@ from veritext.metrics import Readability, ReadabilityResult
|
||||
|
||||
|
||||
class TestReadability:
|
||||
"""Tests for the Readability metric class."""
|
||||
|
||||
@pytest.fixture
|
||||
def readability(self) -> Readability:
|
||||
"""Provide a readability metric instance."""
|
||||
return Readability()
|
||||
|
||||
def test_name(self, readability: Readability) -> None:
|
||||
"""Test that name returns 'readability'."""
|
||||
assert readability.name == "readability"
|
||||
|
||||
def test_requires_reference(self, readability: Readability) -> None:
|
||||
"""Test that readability does NOT require reference text."""
|
||||
assert readability.requires_reference is False
|
||||
|
||||
def test_simple_text(self, readability: Readability) -> None:
|
||||
"""Test readability of simple, easy text."""
|
||||
# Simple children's text - short sentences, simple words
|
||||
text = "The cat sat. The dog ran. I see a bird."
|
||||
result = readability.score(text)
|
||||
@@ -26,10 +32,11 @@ class TestReadability:
|
||||
assert result.flesch_reading_ease > 80.0
|
||||
|
||||
def test_complex_text(self, readability: Readability) -> None:
|
||||
"""Test readability of complex, academic text."""
|
||||
# Complex academic text - long sentences, polysyllabic words
|
||||
text = (
|
||||
"The implementation of sophisticated computational methodologies "
|
||||
"necessitates thorough understanding of algorithmic complexity "
|
||||
"necessitates comprehensive understanding of algorithmic complexity "
|
||||
"and architectural considerations."
|
||||
)
|
||||
result = readability.score(text)
|
||||
@@ -39,6 +46,7 @@ class TestReadability:
|
||||
assert result.flesch_reading_ease < 30.0
|
||||
|
||||
def test_medium_text(self, readability: Readability) -> None:
|
||||
"""Test readability of medium-difficulty text."""
|
||||
text = (
|
||||
"The weather today is quite pleasant. "
|
||||
"Many people are enjoying the sunshine in the park. "
|
||||
@@ -51,6 +59,7 @@ class TestReadability:
|
||||
assert 50.0 < result.flesch_reading_ease < 90.0
|
||||
|
||||
def test_single_sentence(self, readability: Readability) -> None:
|
||||
"""Test readability with a single sentence."""
|
||||
text = "The cat sat on the mat."
|
||||
result = readability.score(text)
|
||||
|
||||
@@ -59,6 +68,7 @@ class TestReadability:
|
||||
assert result.flesch_reading_ease is not None
|
||||
|
||||
def test_single_word(self, readability: Readability) -> None:
|
||||
"""Test readability with a single word."""
|
||||
text = "Cat"
|
||||
result = readability.score(text)
|
||||
|
||||
@@ -67,18 +77,21 @@ class TestReadability:
|
||||
assert result.flesch_reading_ease is not None
|
||||
|
||||
def test_empty_text(self, readability: Readability) -> None:
|
||||
"""Test that empty text returns zero scores."""
|
||||
result = readability.score("")
|
||||
|
||||
assert result.flesch_kincaid_grade == 0.0
|
||||
assert result.flesch_reading_ease == 0.0
|
||||
|
||||
def test_whitespace_only(self, readability: Readability) -> None:
|
||||
"""Test that whitespace-only text returns zero scores."""
|
||||
result = readability.score(" \t\n ")
|
||||
|
||||
assert result.flesch_kincaid_grade == 0.0
|
||||
assert result.flesch_reading_ease == 0.0
|
||||
|
||||
def test_reference_ignored(self, readability: Readability) -> None:
|
||||
"""Test that reference parameter is ignored."""
|
||||
text = "The cat sat on the mat."
|
||||
|
||||
# Score with no reference
|
||||
@@ -94,6 +107,7 @@ class TestReadability:
|
||||
assert result1.flesch_kincaid_grade == result3.flesch_kincaid_grade
|
||||
|
||||
def test_punctuation_handling(self, readability: Readability) -> None:
|
||||
"""Test that punctuation affects sentence counting."""
|
||||
# Same words, different sentence structure
|
||||
text1 = "The cat sat on the mat" # 1 sentence
|
||||
text2 = "The cat sat. On the mat." # 2 sentences
|
||||
@@ -105,6 +119,7 @@ class TestReadability:
|
||||
assert result1.flesch_kincaid_grade != result2.flesch_kincaid_grade
|
||||
|
||||
def test_question_marks_count_sentences(self, readability: Readability) -> None:
|
||||
"""Test that question marks end sentences."""
|
||||
text = "What is this? It is a test."
|
||||
result = readability.score(text)
|
||||
|
||||
@@ -113,6 +128,7 @@ class TestReadability:
|
||||
assert result.flesch_kincaid_grade is not None
|
||||
|
||||
def test_exclamation_marks_count_sentences(self, readability: Readability) -> None:
|
||||
"""Test that exclamation marks end sentences."""
|
||||
text = "Wow! That is amazing!"
|
||||
result = readability.score(text)
|
||||
|
||||
@@ -120,6 +136,7 @@ class TestReadability:
|
||||
assert result.flesch_kincaid_grade is not None
|
||||
|
||||
def test_multiple_punctuation(self, readability: Readability) -> None:
|
||||
"""Test handling of multiple punctuation marks."""
|
||||
text = "What?! That's crazy... Well then."
|
||||
result = readability.score(text)
|
||||
|
||||
@@ -127,10 +144,12 @@ class TestReadability:
|
||||
assert result.flesch_kincaid_grade is not None
|
||||
|
||||
def test_result_score_property(self, readability: Readability) -> None:
|
||||
"""Test that result.score returns flesch_reading_ease."""
|
||||
result = readability.score("The cat sat on the mat.")
|
||||
assert result.score == result.flesch_reading_ease
|
||||
|
||||
def test_contractions(self, readability: Readability) -> None:
|
||||
"""Test handling of contractions."""
|
||||
text = "I'm going to the store. It's not far away."
|
||||
result = readability.score(text)
|
||||
|
||||
@@ -140,11 +159,15 @@ class TestReadability:
|
||||
|
||||
|
||||
class TestReadabilityBatch:
|
||||
"""Tests for readability batch scoring."""
|
||||
|
||||
@pytest.fixture
|
||||
def readability(self) -> Readability:
|
||||
"""Provide a readability metric instance."""
|
||||
return Readability()
|
||||
|
||||
def test_batch_score_basic(self, readability: Readability) -> None:
|
||||
"""Test basic batch scoring."""
|
||||
candidates = [
|
||||
"The cat sat on the mat.",
|
||||
"A dog ran through the park.",
|
||||
@@ -155,6 +178,7 @@ class TestReadabilityBatch:
|
||||
assert len(result.results) == 2
|
||||
|
||||
def test_batch_score_statistics(self, readability: Readability) -> None:
|
||||
"""Test that batch scoring computes statistics."""
|
||||
candidates = [
|
||||
"Cat sat.", # Very simple
|
||||
"The implementation of sophisticated methodologies requires expertise.",
|
||||
@@ -172,6 +196,7 @@ class TestReadabilityBatch:
|
||||
)
|
||||
|
||||
def test_batch_score_percentiles(self, readability: Readability) -> None:
|
||||
"""Test that batch scoring computes percentiles."""
|
||||
candidates = ["a", "b", "c", "d", "e"]
|
||||
result = readability.batch_score(candidates)
|
||||
|
||||
@@ -182,6 +207,7 @@ class TestReadabilityBatch:
|
||||
assert 95 in stats.percentiles
|
||||
|
||||
def test_batch_score_references_ignored(self, readability: Readability) -> None:
|
||||
"""Test that batch scoring ignores references."""
|
||||
candidates = ["The cat sat.", "A dog ran."]
|
||||
|
||||
result1 = readability.batch_score(candidates)
|
||||
@@ -193,12 +219,16 @@ class TestReadabilityBatch:
|
||||
)
|
||||
|
||||
def test_batch_score_empty_list_raises(self, readability: Readability) -> None:
|
||||
"""Test that empty candidate list raises ValueError."""
|
||||
with pytest.raises(ValueError, match="empty"):
|
||||
readability.batch_score([])
|
||||
|
||||
|
||||
class TestReadabilityResult:
|
||||
"""Tests for ReadabilityResult type."""
|
||||
|
||||
def test_frozen(self) -> None:
|
||||
"""Test that ReadabilityResult is frozen."""
|
||||
from pydantic import ValidationError
|
||||
|
||||
result = ReadabilityResult(flesch_kincaid_grade=5.0, flesch_reading_ease=70.0)
|
||||
@@ -206,21 +236,27 @@ class TestReadabilityResult:
|
||||
result.flesch_kincaid_grade = 6.0 # type: ignore[misc]
|
||||
|
||||
def test_values(self) -> None:
|
||||
"""Test that values are stored correctly."""
|
||||
result = ReadabilityResult(flesch_kincaid_grade=8.5, flesch_reading_ease=65.0)
|
||||
assert result.flesch_kincaid_grade == 8.5
|
||||
assert result.flesch_reading_ease == 65.0
|
||||
|
||||
def test_score_property(self) -> None:
|
||||
"""Test that score property returns flesch_reading_ease."""
|
||||
result = ReadabilityResult(flesch_kincaid_grade=8.5, flesch_reading_ease=65.0)
|
||||
assert result.score == 65.0
|
||||
|
||||
|
||||
class TestSyllableCounting:
|
||||
"""Tests for syllable counting heuristics."""
|
||||
|
||||
@pytest.fixture
|
||||
def readability(self) -> Readability:
|
||||
"""Provide a readability metric instance."""
|
||||
return Readability()
|
||||
|
||||
def test_monosyllabic_words(self, readability: Readability) -> None:
|
||||
"""Test that monosyllabic words don't inflate scores."""
|
||||
# All one-syllable words
|
||||
text = "The cat sat on the mat."
|
||||
result = readability.score(text)
|
||||
@@ -229,6 +265,7 @@ class TestSyllableCounting:
|
||||
assert result.flesch_reading_ease > 90.0
|
||||
|
||||
def test_polysyllabic_words(self, readability: Readability) -> None:
|
||||
"""Test that polysyllabic words affect scores."""
|
||||
# Words with multiple syllables
|
||||
text = "International communication facilitates understanding."
|
||||
result = readability.score(text)
|
||||
|
||||
@@ -6,17 +6,23 @@ from veritext.metrics import Rouge, RougeResult, RougeScore
|
||||
|
||||
|
||||
class TestRouge:
|
||||
"""Tests for the Rouge metric class."""
|
||||
|
||||
@pytest.fixture
|
||||
def rouge(self) -> Rouge:
|
||||
"""Provide a ROUGE metric instance."""
|
||||
return Rouge()
|
||||
|
||||
def test_name(self, rouge: Rouge) -> None:
|
||||
"""Test that name returns 'rouge'."""
|
||||
assert rouge.name == "rouge"
|
||||
|
||||
def test_requires_reference(self, rouge: Rouge) -> None:
|
||||
"""Test that ROUGE requires reference text."""
|
||||
assert rouge.requires_reference is True
|
||||
|
||||
def test_identical_texts(self, rouge: Rouge) -> None:
|
||||
"""Test that identical texts produce perfect scores."""
|
||||
text = "The cat sat on the mat"
|
||||
result = rouge.score(text, text)
|
||||
|
||||
@@ -27,6 +33,7 @@ class TestRouge:
|
||||
assert result.rouge_l.fmeasure == 1.0
|
||||
|
||||
def test_no_overlap(self, rouge: Rouge) -> None:
|
||||
"""Test that texts with no overlap produce zero scores."""
|
||||
candidate = "apple banana cherry"
|
||||
reference = "dog elephant fox"
|
||||
result = rouge.score(candidate, reference)
|
||||
@@ -38,6 +45,7 @@ class TestRouge:
|
||||
assert result.rouge_l.fmeasure == 0.0
|
||||
|
||||
def test_partial_overlap_rouge1(self, rouge: Rouge) -> None:
|
||||
"""Test ROUGE-1 with partial overlap."""
|
||||
candidate = "the cat sat"
|
||||
reference = "the dog sat"
|
||||
result = rouge.score(candidate, reference)
|
||||
@@ -49,6 +57,7 @@ class TestRouge:
|
||||
assert abs(result.rouge1.recall - 2 / 3) < 1e-10
|
||||
|
||||
def test_partial_overlap_rouge2(self, rouge: Rouge) -> None:
|
||||
"""Test ROUGE-2 (bigram) with partial overlap."""
|
||||
candidate = "the cat sat on the mat"
|
||||
reference = "the cat lay on the mat"
|
||||
result = rouge.score(candidate, reference)
|
||||
@@ -61,6 +70,7 @@ class TestRouge:
|
||||
assert abs(result.rouge2.recall - 3 / 5) < 1e-10
|
||||
|
||||
def test_rouge_l_basic(self, rouge: Rouge) -> None:
|
||||
"""Test ROUGE-L (LCS) computation."""
|
||||
candidate = "the cat sat on the mat"
|
||||
reference = "the cat sat"
|
||||
result = rouge.score(candidate, reference)
|
||||
@@ -71,6 +81,7 @@ class TestRouge:
|
||||
assert result.rouge_l.recall == 1.0
|
||||
|
||||
def test_rouge_l_non_contiguous(self, rouge: Rouge) -> None:
|
||||
"""Test ROUGE-L with non-contiguous LCS."""
|
||||
candidate = "the big cat sat"
|
||||
reference = "the cat sat"
|
||||
result = rouge.score(candidate, reference)
|
||||
@@ -81,6 +92,7 @@ class TestRouge:
|
||||
assert result.rouge_l.recall == 1.0
|
||||
|
||||
def test_precision_vs_recall(self, rouge: Rouge) -> None:
|
||||
"""Test that precision and recall differ appropriately."""
|
||||
# Short candidate, long reference
|
||||
candidate = "the cat"
|
||||
reference = "the cat sat on the mat"
|
||||
@@ -92,6 +104,7 @@ class TestRouge:
|
||||
assert result.rouge1.recall < 1.0
|
||||
|
||||
def test_empty_candidate(self, rouge: Rouge) -> None:
|
||||
"""Test that empty candidate returns zero scores."""
|
||||
result = rouge.score("", "The cat sat")
|
||||
|
||||
assert result.rouge1.fmeasure == 0.0
|
||||
@@ -99,20 +112,24 @@ class TestRouge:
|
||||
assert result.rouge_l.fmeasure == 0.0
|
||||
|
||||
def test_whitespace_only_candidate(self, rouge: Rouge) -> None:
|
||||
"""Test that whitespace-only candidate returns zero scores."""
|
||||
result = rouge.score(" \t\n ", "The cat sat")
|
||||
|
||||
assert result.rouge1.fmeasure == 0.0
|
||||
assert result.rouge_l.fmeasure == 0.0
|
||||
|
||||
def test_empty_reference_raises(self, rouge: Rouge) -> None:
|
||||
"""Test that empty reference raises ValueError."""
|
||||
with pytest.raises(ValueError, match="cannot be empty"):
|
||||
rouge.score("The cat sat", "")
|
||||
|
||||
def test_none_reference_raises(self, rouge: Rouge) -> None:
|
||||
"""Test that None reference raises ValueError."""
|
||||
with pytest.raises(ValueError, match="requires reference"):
|
||||
rouge.score("The cat sat", None)
|
||||
|
||||
def test_multiple_references_uses_max(self, rouge: Rouge) -> None:
|
||||
"""Test that multiple references use max scores."""
|
||||
candidate = "the cat sat on the mat"
|
||||
references = [
|
||||
"a dog ran across the room", # Low overlap
|
||||
@@ -125,6 +142,7 @@ class TestRouge:
|
||||
assert result.rouge_l.fmeasure == 1.0
|
||||
|
||||
def test_multiple_references_partial(self, rouge: Rouge) -> None:
|
||||
"""Test multiple references with partial matches."""
|
||||
candidate = "the quick brown fox"
|
||||
references = [
|
||||
"the fast brown fox", # 3/4 match
|
||||
@@ -136,19 +154,23 @@ class TestRouge:
|
||||
assert result.rouge1.fmeasure > 0.0
|
||||
|
||||
def test_result_score_property(self, rouge: Rouge) -> None:
|
||||
"""Test that result.score returns rouge_l.fmeasure."""
|
||||
result = rouge.score("The cat sat", "The cat sat")
|
||||
assert result.score == result.rouge_l.fmeasure
|
||||
|
||||
def test_case_insensitivity(self, rouge: Rouge) -> None:
|
||||
"""Test that ROUGE is case insensitive by default."""
|
||||
result = rouge.score("THE CAT SAT", "the cat sat")
|
||||
assert result.rouge1.fmeasure == 1.0
|
||||
assert result.rouge_l.fmeasure == 1.0
|
||||
|
||||
def test_punctuation_ignored(self, rouge: Rouge) -> None:
|
||||
"""Test that punctuation is ignored by default."""
|
||||
result = rouge.score("The cat sat.", "The cat sat!")
|
||||
assert result.rouge1.fmeasure == 1.0
|
||||
|
||||
def test_single_word(self, rouge: Rouge) -> None:
|
||||
"""Test ROUGE with single word texts."""
|
||||
result = rouge.score("cat", "cat")
|
||||
|
||||
assert result.rouge1.fmeasure == 1.0
|
||||
@@ -157,6 +179,7 @@ class TestRouge:
|
||||
assert result.rouge_l.fmeasure == 1.0
|
||||
|
||||
def test_fmeasure_calculation(self, rouge: Rouge) -> None:
|
||||
"""Test that F-measure is calculated correctly."""
|
||||
# Create a case where P != R
|
||||
candidate = "the cat sat on"
|
||||
reference = "the cat"
|
||||
@@ -169,11 +192,15 @@ class TestRouge:
|
||||
|
||||
|
||||
class TestRougeBatch:
|
||||
"""Tests for ROUGE batch scoring."""
|
||||
|
||||
@pytest.fixture
|
||||
def rouge(self) -> Rouge:
|
||||
"""Provide a ROUGE metric instance."""
|
||||
return Rouge()
|
||||
|
||||
def test_batch_score_basic(self, rouge: Rouge) -> None:
|
||||
"""Test basic batch scoring."""
|
||||
candidates = ["The cat sat", "A dog runs"]
|
||||
references = ["The cat sat", "A dog runs"]
|
||||
result = rouge.batch_score(candidates, references)
|
||||
@@ -183,6 +210,7 @@ class TestRougeBatch:
|
||||
assert all(r.rouge_l.fmeasure == 1.0 for r in result.results)
|
||||
|
||||
def test_batch_score_statistics(self, rouge: Rouge) -> None:
|
||||
"""Test that batch scoring computes statistics."""
|
||||
candidates = ["The cat sat", "Completely different words"]
|
||||
references = ["The cat sat", "The cat sat"]
|
||||
result = rouge.batch_score(candidates, references)
|
||||
@@ -199,6 +227,7 @@ class TestRougeBatch:
|
||||
assert result.results[1].rouge1.fmeasure == 0.0
|
||||
|
||||
def test_batch_score_percentiles(self, rouge: Rouge) -> None:
|
||||
"""Test that batch scoring computes percentiles."""
|
||||
candidates = ["a", "b", "c", "d", "e"]
|
||||
references = ["a", "b", "c", "d", "e"]
|
||||
result = rouge.batch_score(candidates, references)
|
||||
@@ -210,14 +239,17 @@ class TestRougeBatch:
|
||||
assert 95 in stats.percentiles
|
||||
|
||||
def test_batch_score_none_references_raises(self, rouge: Rouge) -> None:
|
||||
"""Test that batch scoring raises for None references."""
|
||||
with pytest.raises(ValueError, match="requires reference"):
|
||||
rouge.batch_score(["text"], None)
|
||||
|
||||
def test_batch_score_length_mismatch_raises(self, rouge: Rouge) -> None:
|
||||
"""Test that batch scoring raises for mismatched lengths."""
|
||||
with pytest.raises(ValueError, match="must match"):
|
||||
rouge.batch_score(["a", "b"], ["a"])
|
||||
|
||||
def test_batch_score_multi_refs(self, rouge: Rouge) -> None:
|
||||
def test_batch_score_with_multiple_references(self, rouge: Rouge) -> None:
|
||||
"""Test batch scoring with multiple references per candidate."""
|
||||
candidates = [
|
||||
"The cat sat on the mat",
|
||||
"A quick brown fox",
|
||||
@@ -235,7 +267,10 @@ class TestRougeBatch:
|
||||
|
||||
|
||||
class TestRougeResult:
|
||||
"""Tests for RougeResult and RougeScore types."""
|
||||
|
||||
def test_rouge_score_frozen(self) -> None:
|
||||
"""Test that RougeScore is frozen."""
|
||||
from pydantic import ValidationError
|
||||
|
||||
score = RougeScore(precision=0.5, recall=0.6, fmeasure=0.55)
|
||||
@@ -243,6 +278,7 @@ class TestRougeResult:
|
||||
score.precision = 0.7 # type: ignore[misc]
|
||||
|
||||
def test_rouge_result_frozen(self) -> None:
|
||||
"""Test that RougeResult is frozen."""
|
||||
from pydantic import ValidationError
|
||||
|
||||
score = RougeScore(precision=0.5, recall=0.6, fmeasure=0.55)
|
||||
@@ -251,6 +287,7 @@ class TestRougeResult:
|
||||
result.rouge1 = score # type: ignore[misc]
|
||||
|
||||
def test_score_property(self) -> None:
|
||||
"""Test that score property returns rouge_l.fmeasure."""
|
||||
r1 = RougeScore(precision=0.9, recall=0.9, fmeasure=0.9)
|
||||
r2 = RougeScore(precision=0.8, recall=0.8, fmeasure=0.8)
|
||||
rl = RougeScore(precision=0.7, recall=0.7, fmeasure=0.7)
|
||||
|
||||
@@ -12,11 +12,13 @@ pytest_plugins = ["pytester"]
|
||||
|
||||
@pytest.fixture
|
||||
def text_validator() -> ValidatorFactory:
|
||||
"""Provide a factory for building validators."""
|
||||
return ValidatorFactory()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def validation_context() -> type:
|
||||
"""Provide a factory for creating ValidationContext objects."""
|
||||
from typing import Any
|
||||
|
||||
from veritext.core.types import ValidationContext
|
||||
|
||||
@@ -6,17 +6,22 @@ from veritext.pytest_plugin import validate_text
|
||||
|
||||
|
||||
class TestValidateTextBasicValidation:
|
||||
def test_valid_length(self) -> None:
|
||||
"""Test basic validation scenarios."""
|
||||
|
||||
def test_passes_with_valid_length(self) -> None:
|
||||
"""Test validation passes when length constraints are met."""
|
||||
text = "The quick brown fox jumps over the lazy dog."
|
||||
validate_text(text, min_length=10, max_length=100)
|
||||
|
||||
def test_too_short(self) -> None:
|
||||
def test_fails_when_too_short(self) -> None:
|
||||
"""Test validation fails when text is below minimum length."""
|
||||
text = "Short."
|
||||
with pytest.raises(AssertionError) as exc_info:
|
||||
validate_text(text, min_length=50)
|
||||
assert "length" in str(exc_info.value).lower()
|
||||
|
||||
def test_too_long(self) -> None:
|
||||
def test_fails_when_too_long(self) -> None:
|
||||
"""Test validation fails when text exceeds maximum length."""
|
||||
text = "A" * 100
|
||||
with pytest.raises(AssertionError) as exc_info:
|
||||
validate_text(text, max_length=50)
|
||||
@@ -24,14 +29,18 @@ class TestValidateTextBasicValidation:
|
||||
|
||||
|
||||
class TestValidateTextReadability:
|
||||
def test_simple_text_passes(self) -> None:
|
||||
"""Test readability validation."""
|
||||
|
||||
def test_passes_with_simple_text(self) -> None:
|
||||
"""Test validation passes for simple, readable text."""
|
||||
text = "The cat sat on the mat. It was a nice day."
|
||||
validate_text(text, max_reading_grade=10.0)
|
||||
|
||||
def test_complex_text_fails(self) -> None:
|
||||
def test_fails_with_complex_text(self) -> None:
|
||||
"""Test validation fails for overly complex text."""
|
||||
text = (
|
||||
"The implementation of sophisticated metacognitive strategies "
|
||||
"necessitates the thorough understanding of epistemological "
|
||||
"necessitates the comprehensive understanding of epistemological "
|
||||
"frameworks and their corresponding methodological implications."
|
||||
)
|
||||
with pytest.raises(AssertionError) as exc_info:
|
||||
@@ -40,21 +49,27 @@ class TestValidateTextReadability:
|
||||
|
||||
|
||||
class TestValidateTextPatterns:
|
||||
def test_contains_pattern(self) -> None:
|
||||
"""Test pattern matching validation."""
|
||||
|
||||
def test_passes_when_contains_pattern(self) -> None:
|
||||
"""Test validation passes when required pattern is present."""
|
||||
text = "Please contact support@example.com for assistance."
|
||||
validate_text(text, must_contain=["support@example.com"])
|
||||
|
||||
def test_missing_pattern(self) -> None:
|
||||
def test_fails_when_missing_required_pattern(self) -> None:
|
||||
"""Test validation fails when required pattern is missing."""
|
||||
text = "Please contact us for assistance."
|
||||
with pytest.raises(AssertionError) as exc_info:
|
||||
validate_text(text, must_contain=["@example.com"])
|
||||
assert "contains" in str(exc_info.value).lower()
|
||||
|
||||
def test_excludes_pattern(self) -> None:
|
||||
def test_passes_when_excludes_pattern(self) -> None:
|
||||
"""Test validation passes when forbidden pattern is absent."""
|
||||
text = "The report is complete and reviewed."
|
||||
validate_text(text, must_exclude=["TODO", "FIXME"])
|
||||
|
||||
def test_forbidden_pattern(self) -> None:
|
||||
def test_fails_when_contains_forbidden_pattern(self) -> None:
|
||||
"""Test validation fails when forbidden pattern is present."""
|
||||
text = "The report is almost done. TODO: add conclusion."
|
||||
with pytest.raises(AssertionError) as exc_info:
|
||||
validate_text(text, must_exclude=["TODO"])
|
||||
@@ -62,24 +77,30 @@ class TestValidateTextPatterns:
|
||||
|
||||
|
||||
class TestValidateTextComparisonMetrics:
|
||||
def test_high_bleu_passes(self) -> None:
|
||||
"""Test comparison-based validation (BLEU, ROUGE)."""
|
||||
|
||||
def test_passes_with_high_bleu_score(self) -> None:
|
||||
"""Test validation passes when BLEU score meets threshold."""
|
||||
reference = "The quick brown fox jumps over the lazy dog."
|
||||
text = "The quick brown fox jumps over the lazy dog."
|
||||
validate_text(text, reference=reference, min_bleu=0.9)
|
||||
|
||||
def test_low_bleu_fails(self) -> None:
|
||||
def test_fails_with_low_bleu_score(self) -> None:
|
||||
"""Test validation fails when BLEU score is below threshold."""
|
||||
reference = "The quick brown fox jumps over the lazy dog."
|
||||
text = "A slow red cat sleeps under the active mouse."
|
||||
with pytest.raises(AssertionError) as exc_info:
|
||||
validate_text(text, reference=reference, min_bleu=0.5)
|
||||
assert "bleu" in str(exc_info.value).lower()
|
||||
|
||||
def test_high_rouge_passes(self) -> None:
|
||||
def test_passes_with_high_rouge_score(self) -> None:
|
||||
"""Test validation passes when ROUGE score meets threshold."""
|
||||
reference = "Machine learning models require extensive training data."
|
||||
text = "Machine learning models need extensive training data."
|
||||
validate_text(text, reference=reference, min_rouge=0.5)
|
||||
|
||||
def test_low_rouge_fails(self) -> None:
|
||||
def test_fails_with_low_rouge_score(self) -> None:
|
||||
"""Test validation fails when ROUGE score is below threshold."""
|
||||
reference = "The algorithm processes input data efficiently."
|
||||
text = "Cats enjoy sleeping in sunny spots."
|
||||
with pytest.raises(AssertionError) as exc_info:
|
||||
@@ -88,25 +109,34 @@ class TestValidateTextComparisonMetrics:
|
||||
|
||||
|
||||
class TestValidateTextErrorHandling:
|
||||
def test_requires_criteria(self) -> None:
|
||||
"""Test error handling and edge cases."""
|
||||
|
||||
def test_raises_value_error_when_no_criteria(self) -> None:
|
||||
"""Test that ValueError is raised when no validation criteria provided."""
|
||||
with pytest.raises(ValueError, match="At least one validation criterion"):
|
||||
validate_text("Some text")
|
||||
|
||||
def test_bleu_requires_reference(self) -> None:
|
||||
def test_raises_value_error_when_bleu_without_reference(self) -> None:
|
||||
"""Test that ValueError is raised when BLEU requested without reference."""
|
||||
with pytest.raises(ValueError, match="Reference text required"):
|
||||
validate_text("Some text", min_bleu=0.5)
|
||||
|
||||
def test_rouge_requires_reference(self) -> None:
|
||||
def test_raises_value_error_when_rouge_without_reference(self) -> None:
|
||||
"""Test that ValueError is raised when ROUGE requested without reference."""
|
||||
with pytest.raises(ValueError, match="Reference text required"):
|
||||
validate_text("Some text", min_rouge=0.5)
|
||||
|
||||
def test_semantic_requires_reference(self) -> None:
|
||||
def test_raises_value_error_when_semantic_without_reference(self) -> None:
|
||||
"""Test that ValueError is raised for semantic without reference."""
|
||||
with pytest.raises(ValueError, match="Reference text required"):
|
||||
validate_text("Some text", min_semantic=0.5)
|
||||
|
||||
|
||||
class TestValidateTextMultipleCriteria:
|
||||
def test_all_criteria_pass(self) -> None:
|
||||
"""Test validation with multiple criteria combined."""
|
||||
|
||||
def test_passes_all_criteria(self) -> None:
|
||||
"""Test validation passes when all criteria are met."""
|
||||
reference = "The quick brown fox jumps over the lazy dog."
|
||||
text = "The quick brown fox jumps over the lazy dog."
|
||||
validate_text(
|
||||
@@ -117,7 +147,8 @@ class TestValidateTextMultipleCriteria:
|
||||
max_length=100,
|
||||
)
|
||||
|
||||
def test_one_failure_fails_all(self) -> None:
|
||||
def test_fails_when_one_criterion_fails(self) -> None:
|
||||
"""Test validation fails when any criterion fails."""
|
||||
reference = "The quick brown fox jumps over the lazy dog."
|
||||
text = "The quick brown fox jumps over the lazy dog."
|
||||
with pytest.raises(AssertionError):
|
||||
@@ -130,13 +161,17 @@ class TestValidateTextMultipleCriteria:
|
||||
|
||||
|
||||
class TestValidateTextFailureMessage:
|
||||
def test_includes_text_preview(self) -> None:
|
||||
"""Test failure message formatting."""
|
||||
|
||||
def test_failure_message_includes_text_preview(self) -> None:
|
||||
"""Test that failure message includes preview of the text."""
|
||||
text = "Short text"
|
||||
with pytest.raises(AssertionError) as exc_info:
|
||||
validate_text(text, min_length=100)
|
||||
assert "Short text" in str(exc_info.value)
|
||||
|
||||
def test_truncates_long_text(self) -> None:
|
||||
def test_failure_message_truncates_long_text(self) -> None:
|
||||
"""Test that long text is truncated in failure message."""
|
||||
text = "A" * 200
|
||||
with pytest.raises(AssertionError) as exc_info:
|
||||
validate_text(text, max_length=50)
|
||||
@@ -144,7 +179,8 @@ class TestValidateTextFailureMessage:
|
||||
assert "..." in message
|
||||
assert "A" * 200 not in message
|
||||
|
||||
def test_includes_check_details(self) -> None:
|
||||
def test_failure_message_includes_check_details(self) -> None:
|
||||
"""Test that failure message includes check name and details."""
|
||||
text = "Short"
|
||||
with pytest.raises(AssertionError) as exc_info:
|
||||
validate_text(text, min_length=100)
|
||||
@@ -154,7 +190,10 @@ class TestValidateTextFailureMessage:
|
||||
|
||||
|
||||
class TestValidateTextListReference:
|
||||
def test_bleu_multi_reference(self) -> None:
|
||||
"""Test validation with list of reference texts."""
|
||||
|
||||
def test_bleu_with_multiple_references(self) -> None:
|
||||
"""Test BLEU validation accepts multiple reference texts."""
|
||||
references = [
|
||||
"The quick brown fox jumps over the lazy dog.",
|
||||
"A fast brown fox leaps over a sleepy dog.",
|
||||
@@ -162,7 +201,8 @@ class TestValidateTextListReference:
|
||||
text = "The quick brown fox jumps over the lazy dog."
|
||||
validate_text(text, reference=references, min_bleu=0.9)
|
||||
|
||||
def test_rouge_multi_reference(self) -> None:
|
||||
def test_rouge_with_multiple_references(self) -> None:
|
||||
"""Test ROUGE validation accepts multiple reference texts."""
|
||||
references = [
|
||||
"Machine learning requires data.",
|
||||
"ML models need training data.",
|
||||
|
||||
@@ -6,7 +6,10 @@ from veritext.validators import bleu, length
|
||||
|
||||
|
||||
class TestValidatorFactory:
|
||||
"""Test the ValidatorFactory class."""
|
||||
|
||||
def test_creates_validator_from_checks(self) -> None:
|
||||
"""Test that factory creates a callable validator."""
|
||||
factory = ValidatorFactory()
|
||||
validate = factory(checks=[length(min_chars=5)])
|
||||
|
||||
@@ -14,6 +17,7 @@ class TestValidatorFactory:
|
||||
assert result.passed
|
||||
|
||||
def test_validator_uses_provided_reference(self) -> None:
|
||||
"""Test that factory passes reference to context."""
|
||||
factory = ValidatorFactory()
|
||||
reference = "The quick brown fox."
|
||||
validate = factory(
|
||||
@@ -26,6 +30,7 @@ class TestValidatorFactory:
|
||||
assert result.passed
|
||||
|
||||
def test_validator_returns_validation_result(self) -> None:
|
||||
"""Test that validator returns a ValidationResult."""
|
||||
factory = ValidatorFactory()
|
||||
validate = factory(checks=[length(min_chars=100)])
|
||||
|
||||
@@ -36,13 +41,17 @@ class TestValidatorFactory:
|
||||
|
||||
|
||||
class TestTextValidatorFixture:
|
||||
"""Test the text_validator fixture."""
|
||||
|
||||
def test_fixture_returns_factory(self, text_validator: ValidatorFactory) -> None:
|
||||
"""Test that fixture provides a ValidatorFactory."""
|
||||
assert isinstance(text_validator, ValidatorFactory)
|
||||
|
||||
def test_fixture_can_create_validators(
|
||||
self,
|
||||
text_validator: ValidatorFactory,
|
||||
) -> None:
|
||||
"""Test that fixture can be used to create validators."""
|
||||
validate = text_validator(checks=[length(min_chars=5, max_chars=50)])
|
||||
|
||||
assert validate("Hello, World!").passed
|
||||
@@ -50,10 +59,13 @@ class TestTextValidatorFixture:
|
||||
|
||||
|
||||
class TestValidationContextFixture:
|
||||
"""Test the validation_context fixture."""
|
||||
|
||||
def test_fixture_creates_context(
|
||||
self,
|
||||
validation_context: type,
|
||||
) -> None:
|
||||
"""Test that fixture creates ValidationContext."""
|
||||
ctx = validation_context(reference="Test reference")
|
||||
assert isinstance(ctx, ValidationContext)
|
||||
assert ctx.reference == "Test reference"
|
||||
@@ -62,6 +74,7 @@ class TestValidationContextFixture:
|
||||
self,
|
||||
validation_context: type,
|
||||
) -> None:
|
||||
"""Test that fixture passes metadata to context."""
|
||||
ctx = validation_context(reference="Test", source="unit_test", version=1)
|
||||
assert ctx.metadata["source"] == "unit_test"
|
||||
assert ctx.metadata["version"] == 1
|
||||
@@ -70,5 +83,6 @@ class TestValidationContextFixture:
|
||||
self,
|
||||
validation_context: type,
|
||||
) -> None:
|
||||
"""Test that fixture allows creating context without reference."""
|
||||
ctx = validation_context()
|
||||
assert ctx.reference is None
|
||||
|
||||
@@ -5,10 +5,16 @@ import pytest
|
||||
|
||||
@pytest.fixture
|
||||
def plugin_pytester(pytester: pytest.Pytester) -> pytest.Pytester:
|
||||
"""Configure pytester to use the veritext plugin.
|
||||
|
||||
Note: The plugin is already loaded via the entry point in pyproject.toml,
|
||||
so no explicit pytest_plugins declaration is needed.
|
||||
"""
|
||||
return pytester
|
||||
|
||||
|
||||
def test_plugin_registers_marker(plugin_pytester: pytest.Pytester) -> None:
|
||||
"""Test that the text_validation marker is registered."""
|
||||
plugin_pytester.makepyfile(
|
||||
"""
|
||||
import pytest
|
||||
@@ -24,6 +30,7 @@ def test_plugin_registers_marker(plugin_pytester: pytest.Pytester) -> None:
|
||||
|
||||
|
||||
def test_marker_can_be_used(plugin_pytester: pytest.Pytester) -> None:
|
||||
"""Test that the text_validation marker can filter tests."""
|
||||
plugin_pytester.makepyfile(
|
||||
"""
|
||||
import pytest
|
||||
@@ -42,6 +49,7 @@ def test_marker_can_be_used(plugin_pytester: pytest.Pytester) -> None:
|
||||
|
||||
|
||||
def test_validate_text_is_importable(plugin_pytester: pytest.Pytester) -> None:
|
||||
"""Test that validate_text can be imported from the plugin."""
|
||||
plugin_pytester.makepyfile(
|
||||
"""
|
||||
from veritext.pytest_plugin import validate_text
|
||||
@@ -55,6 +63,7 @@ def test_validate_text_is_importable(plugin_pytester: pytest.Pytester) -> None:
|
||||
|
||||
|
||||
def test_validate_text_works_in_tests(plugin_pytester: pytest.Pytester) -> None:
|
||||
"""Test that validate_text can be used in test functions."""
|
||||
plugin_pytester.makepyfile(
|
||||
"""
|
||||
from veritext.pytest_plugin import validate_text
|
||||
@@ -72,6 +81,7 @@ def test_validate_text_works_in_tests(plugin_pytester: pytest.Pytester) -> None:
|
||||
|
||||
|
||||
def test_validate_text_failure_in_tests(plugin_pytester: pytest.Pytester) -> None:
|
||||
"""Test that validate_text failures are reported properly."""
|
||||
plugin_pytester.makepyfile(
|
||||
"""
|
||||
from veritext.pytest_plugin import validate_text
|
||||
|
||||
@@ -10,17 +10,23 @@ from veritext.semantic import SemanticSimilarity
|
||||
|
||||
|
||||
class TestSemanticSimilarity:
|
||||
"""Tests for the SemanticSimilarity metric class."""
|
||||
|
||||
@pytest.fixture
|
||||
def semantic(self) -> SemanticSimilarity:
|
||||
"""Provide a SemanticSimilarity metric instance."""
|
||||
return SemanticSimilarity()
|
||||
|
||||
def test_name(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that name returns 'semantic'."""
|
||||
assert semantic.name == "semantic"
|
||||
|
||||
def test_requires_reference(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that semantic similarity requires reference text."""
|
||||
assert semantic.requires_reference is True
|
||||
|
||||
def test_identical_texts(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that identical texts produce high similarity."""
|
||||
text = "The cat sat on the mat"
|
||||
result = semantic.score(text, text)
|
||||
|
||||
@@ -29,6 +35,7 @@ class TestSemanticSimilarity:
|
||||
assert result.model == "all-MiniLM-L6-v2"
|
||||
|
||||
def test_semantically_similar_texts(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that semantically similar texts have high similarity."""
|
||||
candidate = "The cat sat on the mat"
|
||||
reference = "A feline rested on the rug"
|
||||
result = semantic.score(candidate, reference)
|
||||
@@ -37,6 +44,7 @@ class TestSemanticSimilarity:
|
||||
assert result.similarity > 0.3
|
||||
|
||||
def test_unrelated_texts(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that unrelated texts have low similarity."""
|
||||
candidate = "The quick brown fox"
|
||||
reference = "Quantum physics describes particle behaviour"
|
||||
result = semantic.score(candidate, reference)
|
||||
@@ -45,26 +53,32 @@ class TestSemanticSimilarity:
|
||||
assert result.similarity < 0.5
|
||||
|
||||
def test_empty_candidate(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that empty candidate returns zero similarity."""
|
||||
result = semantic.score("", "The cat sat on the mat")
|
||||
assert result.similarity == 0.0
|
||||
|
||||
def test_whitespace_only_candidate(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that whitespace-only candidate returns zero similarity."""
|
||||
result = semantic.score(" \t\n ", "The cat sat on the mat")
|
||||
assert result.similarity == 0.0
|
||||
|
||||
def test_none_reference_raises(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that None reference raises ValueError."""
|
||||
with pytest.raises(ValueError, match="requires reference"):
|
||||
semantic.score("The cat sat", None)
|
||||
|
||||
def test_empty_reference_raises(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that empty reference raises ValueError."""
|
||||
with pytest.raises(ValueError, match="cannot be empty"):
|
||||
semantic.score("The cat sat", "")
|
||||
|
||||
def test_whitespace_reference_raises(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that whitespace-only reference raises ValueError."""
|
||||
with pytest.raises(ValueError, match="cannot be empty"):
|
||||
semantic.score("The cat sat", " \t\n ")
|
||||
|
||||
def test_multiple_references(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test semantic similarity with multiple references uses max."""
|
||||
candidate = "The cat sat on the mat"
|
||||
references = [
|
||||
"A dog ran through the park",
|
||||
@@ -76,6 +90,7 @@ class TestSemanticSimilarity:
|
||||
assert result.similarity >= 0.99
|
||||
|
||||
def test_multiple_references_takes_max(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that multiple references returns maximum similarity."""
|
||||
candidate = "The cat sat on the mat"
|
||||
references = [
|
||||
"Quantum physics is complex", # Low similarity
|
||||
@@ -87,10 +102,12 @@ class TestSemanticSimilarity:
|
||||
assert result.similarity > 0.3
|
||||
|
||||
def test_result_score_property(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that result.score returns similarity."""
|
||||
result = semantic.score("The cat sat", "The cat sat")
|
||||
assert result.score == result.similarity
|
||||
|
||||
def test_caching_behaviour(self) -> None:
|
||||
"""Test that caching works for repeated texts."""
|
||||
semantic = SemanticSimilarity(cache_embeddings=True)
|
||||
|
||||
# Score same texts multiple times
|
||||
@@ -107,6 +124,7 @@ class TestSemanticSimilarity:
|
||||
assert result3.similarity == result1.similarity
|
||||
|
||||
def test_caching_disabled(self) -> None:
|
||||
"""Test that caching can be disabled."""
|
||||
semantic = SemanticSimilarity(cache_embeddings=False)
|
||||
|
||||
text = "The cat sat on the mat"
|
||||
@@ -120,6 +138,7 @@ class TestSemanticSimilarity:
|
||||
semantic.clear_cache()
|
||||
|
||||
def test_custom_model(self) -> None:
|
||||
"""Test that custom model name is recorded in result."""
|
||||
# Use the same model but verify it's recorded correctly
|
||||
semantic = SemanticSimilarity(model="all-MiniLM-L6-v2")
|
||||
result = semantic.score("Test text", "Test text")
|
||||
@@ -127,11 +146,15 @@ class TestSemanticSimilarity:
|
||||
|
||||
|
||||
class TestSemanticSimilarityBatch:
|
||||
"""Tests for semantic similarity batch scoring."""
|
||||
|
||||
@pytest.fixture
|
||||
def semantic(self) -> SemanticSimilarity:
|
||||
"""Provide a SemanticSimilarity metric instance."""
|
||||
return SemanticSimilarity()
|
||||
|
||||
def test_batch_score_basic(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test basic batch scoring."""
|
||||
candidates = ["The cat sat on the mat", "A quick brown dog runs fast"]
|
||||
references = ["The cat sat on the mat", "A quick brown dog runs fast"]
|
||||
result = semantic.batch_score(candidates, references)
|
||||
@@ -142,6 +165,7 @@ class TestSemanticSimilarityBatch:
|
||||
assert all(r.similarity >= 0.99 for r in result.results)
|
||||
|
||||
def test_batch_score_statistics(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that batch scoring computes statistics."""
|
||||
candidates = ["The cat sat", "Quantum physics is complex"]
|
||||
references = ["The cat sat", "The cat sat"]
|
||||
result = semantic.batch_score(candidates, references)
|
||||
@@ -154,6 +178,7 @@ class TestSemanticSimilarityBatch:
|
||||
assert stats.min <= stats.mean <= stats.max
|
||||
|
||||
def test_batch_score_percentiles(self, semantic: SemanticSimilarity) -> None:
|
||||
"""Test that batch scoring computes percentiles."""
|
||||
candidates = ["a", "b", "c", "d", "e"]
|
||||
references = ["a", "b", "c", "d", "e"]
|
||||
result = semantic.batch_score(candidates, references)
|
||||
@@ -167,18 +192,21 @@ class TestSemanticSimilarityBatch:
|
||||
def test_batch_score_none_references_raises(
|
||||
self, semantic: SemanticSimilarity
|
||||
) -> None:
|
||||
"""Test that batch scoring raises for None references."""
|
||||
with pytest.raises(ValueError, match="requires reference"):
|
||||
semantic.batch_score(["text"], None)
|
||||
|
||||
def test_batch_score_length_mismatch_raises(
|
||||
self, semantic: SemanticSimilarity
|
||||
) -> None:
|
||||
"""Test that batch scoring raises for mismatched lengths."""
|
||||
with pytest.raises(ValueError, match="must match"):
|
||||
semantic.batch_score(["a", "b"], ["a"])
|
||||
|
||||
def test_batch_score_multi_refs(
|
||||
def test_batch_score_with_multiple_references(
|
||||
self, semantic: SemanticSimilarity
|
||||
) -> None:
|
||||
"""Test batch scoring with multiple references per candidate."""
|
||||
candidates = [
|
||||
"The cat sat on the mat",
|
||||
"A quick brown dog runs fast",
|
||||
@@ -196,7 +224,10 @@ class TestSemanticSimilarityBatch:
|
||||
|
||||
|
||||
class TestSemanticResult:
|
||||
"""Tests for SemanticResult type."""
|
||||
|
||||
def test_frozen(self) -> None:
|
||||
"""Test that SemanticResult is frozen."""
|
||||
from pydantic import ValidationError
|
||||
|
||||
result = SemanticResult(similarity=0.85, model="test-model")
|
||||
@@ -204,5 +235,6 @@ class TestSemanticResult:
|
||||
result.similarity = 0.9 # type: ignore[misc]
|
||||
|
||||
def test_score_property(self) -> None:
|
||||
"""Test that score property returns similarity."""
|
||||
result = SemanticResult(similarity=0.75, model="test-model")
|
||||
assert result.score == 0.75
|
||||
|
||||
@@ -8,7 +8,10 @@ from veritext.validators.composite import AllOf, AnyOf
|
||||
|
||||
|
||||
class TestAllOf:
|
||||
"""Tests for AllOf composite validator."""
|
||||
|
||||
def test_all_of_passes_when_all_checks_pass(self) -> None:
|
||||
"""Test that AllOf passes when all checks pass."""
|
||||
validator = AllOf(
|
||||
checks=[
|
||||
length(min_words=2),
|
||||
@@ -23,6 +26,7 @@ class TestAllOf:
|
||||
assert all(c.passed for c in result.checks)
|
||||
|
||||
def test_all_of_fails_when_one_check_fails(self) -> None:
|
||||
"""Test that AllOf fails when any check fails."""
|
||||
validator = AllOf(
|
||||
checks=[
|
||||
length(min_words=2),
|
||||
@@ -37,6 +41,7 @@ class TestAllOf:
|
||||
assert len(result.failed_checks) == 1
|
||||
|
||||
def test_all_of_fails_when_all_checks_fail(self) -> None:
|
||||
"""Test that AllOf fails when all checks fail."""
|
||||
validator = AllOf(
|
||||
checks=[
|
||||
length(min_words=10),
|
||||
@@ -50,6 +55,7 @@ class TestAllOf:
|
||||
assert len(result.failed_checks) == 2
|
||||
|
||||
def test_all_of_with_metric_validators(self) -> None:
|
||||
"""Test AllOf with metric-based validators."""
|
||||
validator = AllOf(
|
||||
checks=[
|
||||
bleu(min_score=0.5),
|
||||
@@ -63,6 +69,7 @@ class TestAllOf:
|
||||
assert len(result.checks) == 2
|
||||
|
||||
def test_all_of_failure_summary(self) -> None:
|
||||
"""Test the failure summary property."""
|
||||
validator = AllOf(
|
||||
checks=[
|
||||
length(min_words=10),
|
||||
@@ -78,20 +85,26 @@ class TestAllOf:
|
||||
assert "contains" in summary
|
||||
|
||||
def test_all_of_raises_on_empty_checks(self) -> None:
|
||||
"""Test that empty checks list raises error."""
|
||||
with pytest.raises(ValueError, match="cannot be empty"):
|
||||
AllOf(checks=[])
|
||||
|
||||
def test_all_of_name_property(self) -> None:
|
||||
"""Test the name property."""
|
||||
validator = AllOf(checks=[length(min_chars=1)])
|
||||
assert validator.name == "all_of"
|
||||
|
||||
def test_all_of_factory_function(self) -> None:
|
||||
"""Test the all_of() factory function."""
|
||||
validator = all_of(checks=[length(min_chars=1)])
|
||||
assert isinstance(validator, AllOf)
|
||||
|
||||
|
||||
class TestAnyOf:
|
||||
"""Tests for AnyOf composite validator."""
|
||||
|
||||
def test_any_of_passes_when_any_check_passes(self) -> None:
|
||||
"""Test that AnyOf passes when any check passes."""
|
||||
validator = AnyOf(
|
||||
checks=[
|
||||
length(min_words=10), # Will fail
|
||||
@@ -107,6 +120,7 @@ class TestAnyOf:
|
||||
assert any(c.passed for c in result.checks)
|
||||
|
||||
def test_any_of_passes_when_all_checks_pass(self) -> None:
|
||||
"""Test that AnyOf passes when all checks pass."""
|
||||
validator = AnyOf(
|
||||
checks=[
|
||||
length(min_words=2),
|
||||
@@ -120,6 +134,7 @@ class TestAnyOf:
|
||||
assert all(c.passed for c in result.checks)
|
||||
|
||||
def test_any_of_fails_when_all_checks_fail(self) -> None:
|
||||
"""Test that AnyOf fails when all checks fail."""
|
||||
validator = AnyOf(
|
||||
checks=[
|
||||
length(min_words=10),
|
||||
@@ -133,6 +148,7 @@ class TestAnyOf:
|
||||
assert not any(c.passed for c in result.checks)
|
||||
|
||||
def test_any_of_with_metric_validators(self) -> None:
|
||||
"""Test AnyOf with metric-based validators."""
|
||||
validator = AnyOf(
|
||||
checks=[
|
||||
bleu(min_score=0.9), # Might fail
|
||||
@@ -145,6 +161,7 @@ class TestAnyOf:
|
||||
assert result.passed is True # Length check passes
|
||||
|
||||
def test_any_of_with_excludes(self) -> None:
|
||||
"""Test AnyOf with excludes validator."""
|
||||
validator = AnyOf(
|
||||
checks=[
|
||||
excludes(patterns=["error"]),
|
||||
@@ -166,13 +183,16 @@ class TestAnyOf:
|
||||
assert result.passed is False
|
||||
|
||||
def test_any_of_raises_on_empty_checks(self) -> None:
|
||||
"""Test that empty checks list raises error."""
|
||||
with pytest.raises(ValueError, match="cannot be empty"):
|
||||
AnyOf(checks=[])
|
||||
|
||||
def test_any_of_name_property(self) -> None:
|
||||
"""Test the name property."""
|
||||
validator = AnyOf(checks=[length(min_chars=1)])
|
||||
assert validator.name == "any_of"
|
||||
|
||||
def test_any_of_factory_function(self) -> None:
|
||||
"""Test the any_of() factory function."""
|
||||
validator = any_of(checks=[length(min_chars=1)])
|
||||
assert isinstance(validator, AnyOf)
|
||||
|
||||
@@ -14,7 +14,10 @@ from veritext.validators.constraint import (
|
||||
|
||||
|
||||
class TestLengthValidator:
|
||||
def test_min_chars_passes(self) -> None:
|
||||
"""Tests for LengthValidator."""
|
||||
|
||||
def test_length_validator_min_chars_passes(self) -> None:
|
||||
"""Test that validator passes when char count meets minimum."""
|
||||
validator = LengthValidator(min_chars=10)
|
||||
context = ValidationContext()
|
||||
result = validator.check("hello world!", context)
|
||||
@@ -23,7 +26,8 @@ class TestLengthValidator:
|
||||
assert result.name == "length"
|
||||
assert result.actual["chars"] == 12
|
||||
|
||||
def test_min_chars_fails(self) -> None:
|
||||
def test_length_validator_min_chars_fails(self) -> None:
|
||||
"""Test that validator fails when char count below minimum."""
|
||||
validator = LengthValidator(min_chars=20)
|
||||
context = ValidationContext()
|
||||
result = validator.check("hello", context)
|
||||
@@ -31,7 +35,8 @@ class TestLengthValidator:
|
||||
assert result.passed is False
|
||||
assert "< min" in result.message
|
||||
|
||||
def test_max_chars_passes(self) -> None:
|
||||
def test_length_validator_max_chars_passes(self) -> None:
|
||||
"""Test that validator passes when char count within maximum."""
|
||||
validator = LengthValidator(max_chars=20)
|
||||
context = ValidationContext()
|
||||
result = validator.check("hello world", context)
|
||||
@@ -39,7 +44,8 @@ class TestLengthValidator:
|
||||
assert result.passed is True
|
||||
assert result.actual["chars"] == 11
|
||||
|
||||
def test_max_chars_fails(self) -> None:
|
||||
def test_length_validator_max_chars_fails(self) -> None:
|
||||
"""Test that validator fails when char count exceeds maximum."""
|
||||
validator = LengthValidator(max_chars=5)
|
||||
context = ValidationContext()
|
||||
result = validator.check("hello world", context)
|
||||
@@ -47,7 +53,8 @@ class TestLengthValidator:
|
||||
assert result.passed is False
|
||||
assert "> max" in result.message
|
||||
|
||||
def test_min_words_passes(self) -> None:
|
||||
def test_length_validator_min_words_passes(self) -> None:
|
||||
"""Test that validator passes when word count meets minimum."""
|
||||
validator = LengthValidator(min_words=3)
|
||||
context = ValidationContext()
|
||||
result = validator.check("the quick brown fox", context)
|
||||
@@ -55,7 +62,8 @@ class TestLengthValidator:
|
||||
assert result.passed is True
|
||||
assert result.actual["words"] == 4
|
||||
|
||||
def test_min_words_fails(self) -> None:
|
||||
def test_length_validator_min_words_fails(self) -> None:
|
||||
"""Test that validator fails when word count below minimum."""
|
||||
validator = LengthValidator(min_words=10)
|
||||
context = ValidationContext()
|
||||
result = validator.check("hello world", context)
|
||||
@@ -63,14 +71,16 @@ class TestLengthValidator:
|
||||
assert result.passed is False
|
||||
assert "words < min" in result.message
|
||||
|
||||
def test_max_words_passes(self) -> None:
|
||||
def test_length_validator_max_words_passes(self) -> None:
|
||||
"""Test that validator passes when word count within maximum."""
|
||||
validator = LengthValidator(max_words=5)
|
||||
context = ValidationContext()
|
||||
result = validator.check("hello world", context)
|
||||
|
||||
assert result.passed is True
|
||||
|
||||
def test_max_words_fails(self) -> None:
|
||||
def test_length_validator_max_words_fails(self) -> None:
|
||||
"""Test that validator fails when word count exceeds maximum."""
|
||||
validator = LengthValidator(max_words=2)
|
||||
context = ValidationContext()
|
||||
result = validator.check("the quick brown fox", context)
|
||||
@@ -78,7 +88,8 @@ class TestLengthValidator:
|
||||
assert result.passed is False
|
||||
assert "words > max" in result.message
|
||||
|
||||
def test_combined_constraints(self) -> None:
|
||||
def test_length_validator_combined_constraints(self) -> None:
|
||||
"""Test validator with multiple constraints."""
|
||||
validator = LengthValidator(
|
||||
min_chars=5, max_chars=50, min_words=2, max_words=10
|
||||
)
|
||||
@@ -91,11 +102,13 @@ class TestLengthValidator:
|
||||
assert "min_words" in result.threshold
|
||||
assert "max_words" in result.threshold
|
||||
|
||||
def test_requires_constraints(self) -> None:
|
||||
def test_length_validator_raises_when_no_constraints(self) -> None:
|
||||
"""Test that validator raises when no constraints provided."""
|
||||
with pytest.raises(InvalidThresholdError, match="At least one"):
|
||||
LengthValidator()
|
||||
|
||||
def test_rejects_negative_values(self) -> None:
|
||||
def test_length_validator_raises_on_negative_values(self) -> None:
|
||||
"""Test that negative constraint values raise error."""
|
||||
with pytest.raises(InvalidThresholdError, match="min_chars must be >= 0"):
|
||||
LengthValidator(min_chars=-1)
|
||||
|
||||
@@ -108,21 +121,26 @@ class TestLengthValidator:
|
||||
with pytest.raises(InvalidThresholdError, match="max_words must be >= 0"):
|
||||
LengthValidator(max_words=-1)
|
||||
|
||||
def test_rejects_inverted_range(self) -> None:
|
||||
def test_length_validator_raises_on_invalid_range(self) -> None:
|
||||
"""Test that min > max raises error."""
|
||||
with pytest.raises(InvalidThresholdError, match="cannot exceed max_chars"):
|
||||
LengthValidator(min_chars=100, max_chars=50)
|
||||
|
||||
with pytest.raises(InvalidThresholdError, match="cannot exceed max_words"):
|
||||
LengthValidator(min_words=20, max_words=5)
|
||||
|
||||
def test_length_factory(self) -> None:
|
||||
def test_length_factory_function(self) -> None:
|
||||
"""Test the length() factory function."""
|
||||
validator = length(min_chars=10, max_words=100)
|
||||
assert isinstance(validator, LengthValidator)
|
||||
assert validator.name == "length"
|
||||
|
||||
|
||||
class TestReadabilityValidator:
|
||||
def test_max_grade_passes(self) -> None:
|
||||
"""Tests for ReadabilityValidator."""
|
||||
|
||||
def test_readability_validator_max_grade_passes(self) -> None:
|
||||
"""Test that validator passes when grade level within maximum."""
|
||||
validator = ReadabilityValidator(max_grade=12.0)
|
||||
context = ValidationContext()
|
||||
# Simple text should have low grade level
|
||||
@@ -132,13 +150,14 @@ class TestReadabilityValidator:
|
||||
assert result.name == "readability"
|
||||
assert "grade" in result.actual
|
||||
|
||||
def test_max_grade_fails(self) -> None:
|
||||
def test_readability_validator_max_grade_fails(self) -> None:
|
||||
"""Test that validator fails when grade level exceeds maximum."""
|
||||
validator = ReadabilityValidator(max_grade=1.0)
|
||||
context = ValidationContext()
|
||||
# Complex text
|
||||
result = validator.check(
|
||||
"The implementation of sophisticated methodologies necessitates "
|
||||
"detailed analytical frameworks for systematic evaluation.",
|
||||
"comprehensive analytical frameworks for systematic evaluation.",
|
||||
context,
|
||||
)
|
||||
|
||||
@@ -146,7 +165,8 @@ class TestReadabilityValidator:
|
||||
assert "grade level" in result.message
|
||||
assert "> max" in result.message
|
||||
|
||||
def test_min_ease_passes(self) -> None:
|
||||
def test_readability_validator_min_ease_passes(self) -> None:
|
||||
"""Test that validator passes when reading ease meets minimum."""
|
||||
validator = ReadabilityValidator(min_ease=30.0)
|
||||
context = ValidationContext()
|
||||
# Simple text should have high reading ease
|
||||
@@ -155,12 +175,13 @@ class TestReadabilityValidator:
|
||||
assert result.passed is True
|
||||
assert "ease" in result.actual
|
||||
|
||||
def test_min_ease_fails(self) -> None:
|
||||
def test_readability_validator_min_ease_fails(self) -> None:
|
||||
"""Test that validator fails when reading ease below minimum."""
|
||||
validator = ReadabilityValidator(min_ease=100.0)
|
||||
context = ValidationContext()
|
||||
result = validator.check(
|
||||
"The implementation of sophisticated methodologies necessitates "
|
||||
"detailed analytical frameworks.",
|
||||
"comprehensive analytical frameworks.",
|
||||
context,
|
||||
)
|
||||
|
||||
@@ -168,7 +189,8 @@ class TestReadabilityValidator:
|
||||
assert "reading ease" in result.message
|
||||
assert "< min" in result.message
|
||||
|
||||
def test_combined_grade_and_ease(self) -> None:
|
||||
def test_readability_validator_combined_constraints(self) -> None:
|
||||
"""Test validator with both grade and ease constraints."""
|
||||
validator = ReadabilityValidator(max_grade=12.0, min_ease=30.0)
|
||||
context = ValidationContext()
|
||||
result = validator.check("The cat sat on the mat.", context)
|
||||
@@ -176,18 +198,23 @@ class TestReadabilityValidator:
|
||||
assert "max_grade" in result.threshold
|
||||
assert "min_ease" in result.threshold
|
||||
|
||||
def test_requires_constraints(self) -> None:
|
||||
def test_readability_validator_raises_when_no_constraints(self) -> None:
|
||||
"""Test that validator raises when no constraints provided."""
|
||||
with pytest.raises(InvalidThresholdError, match="At least one"):
|
||||
ReadabilityValidator()
|
||||
|
||||
def test_readability_factory(self) -> None:
|
||||
def test_readability_factory_function(self) -> None:
|
||||
"""Test the readability() factory function."""
|
||||
validator = readability(max_grade=8.0, min_ease=60.0)
|
||||
assert isinstance(validator, ReadabilityValidator)
|
||||
assert validator.name == "readability"
|
||||
|
||||
|
||||
class TestContainsValidator:
|
||||
def test_finds_all_patterns(self) -> None:
|
||||
"""Tests for ContainsValidator."""
|
||||
|
||||
def test_contains_validator_passes_when_pattern_found(self) -> None:
|
||||
"""Test that validator passes when all patterns are found."""
|
||||
validator = ContainsValidator(patterns=["hello", "world"])
|
||||
context = ValidationContext()
|
||||
result = validator.check("Hello World!", context)
|
||||
@@ -197,7 +224,8 @@ class TestContainsValidator:
|
||||
assert result.actual["found"] == 2
|
||||
assert result.actual["missing"] == []
|
||||
|
||||
def test_fails_on_missing_pattern(self) -> None:
|
||||
def test_contains_validator_fails_when_pattern_missing(self) -> None:
|
||||
"""Test that validator fails when a pattern is missing."""
|
||||
validator = ContainsValidator(patterns=["hello", "goodbye"])
|
||||
context = ValidationContext()
|
||||
result = validator.check("Hello World!", context)
|
||||
@@ -206,43 +234,52 @@ class TestContainsValidator:
|
||||
assert "goodbye" in result.actual["missing"]
|
||||
assert "missing" in result.message
|
||||
|
||||
def test_case_insensitive_default(self) -> None:
|
||||
def test_contains_validator_case_insensitive_by_default(self) -> None:
|
||||
"""Test that matching is case-insensitive by default."""
|
||||
validator = ContainsValidator(patterns=["HELLO"])
|
||||
context = ValidationContext()
|
||||
result = validator.check("hello world", context)
|
||||
|
||||
assert result.passed is True
|
||||
|
||||
def test_case_sensitive(self) -> None:
|
||||
def test_contains_validator_case_sensitive(self) -> None:
|
||||
"""Test case-sensitive matching."""
|
||||
validator = ContainsValidator(patterns=["HELLO"], case_sensitive=True)
|
||||
context = ValidationContext()
|
||||
result = validator.check("hello world", context)
|
||||
|
||||
assert result.passed is False
|
||||
|
||||
def test_regex_patterns(self) -> None:
|
||||
def test_contains_validator_regex_patterns(self) -> None:
|
||||
"""Test regex pattern matching."""
|
||||
validator = ContainsValidator(patterns=[r"\d{3}-\d{4}"])
|
||||
context = ValidationContext()
|
||||
result = validator.check("Call 555-1234 for info", context)
|
||||
|
||||
assert result.passed is True
|
||||
|
||||
def test_rejects_empty_patterns(self) -> None:
|
||||
def test_contains_validator_raises_on_empty_patterns(self) -> None:
|
||||
"""Test that empty patterns list raises error."""
|
||||
with pytest.raises(InvalidThresholdError, match="cannot be empty"):
|
||||
ContainsValidator(patterns=[])
|
||||
|
||||
def test_rejects_invalid_regex(self) -> None:
|
||||
def test_contains_validator_raises_on_invalid_regex(self) -> None:
|
||||
"""Test that invalid regex pattern raises error at init time."""
|
||||
with pytest.raises(InvalidThresholdError, match="Invalid regex"):
|
||||
ContainsValidator(patterns=[r"[invalid"])
|
||||
|
||||
def test_contains_factory(self) -> None:
|
||||
def test_contains_factory_function(self) -> None:
|
||||
"""Test the contains() factory function."""
|
||||
validator = contains(patterns=["test"], case_sensitive=True)
|
||||
assert isinstance(validator, ContainsValidator)
|
||||
assert validator.name == "contains"
|
||||
|
||||
|
||||
class TestExcludesValidator:
|
||||
def test_passes_when_absent(self) -> None:
|
||||
"""Tests for ExcludesValidator."""
|
||||
|
||||
def test_excludes_validator_passes_when_pattern_absent(self) -> None:
|
||||
"""Test that validator passes when all patterns are absent."""
|
||||
validator = ExcludesValidator(patterns=["bad", "forbidden"])
|
||||
context = ValidationContext()
|
||||
result = validator.check("This is good text.", context)
|
||||
@@ -251,7 +288,8 @@ class TestExcludesValidator:
|
||||
assert result.name == "excludes"
|
||||
assert result.actual["found"] == []
|
||||
|
||||
def test_fails_when_found(self) -> None:
|
||||
def test_excludes_validator_fails_when_pattern_found(self) -> None:
|
||||
"""Test that validator fails when a forbidden pattern is found."""
|
||||
validator = ExcludesValidator(patterns=["bad", "forbidden"])
|
||||
context = ValidationContext()
|
||||
result = validator.check("This is bad text.", context)
|
||||
@@ -260,21 +298,24 @@ class TestExcludesValidator:
|
||||
assert "bad" in result.actual["found"]
|
||||
assert "forbidden" in result.message
|
||||
|
||||
def test_case_insensitive_default(self) -> None:
|
||||
def test_excludes_validator_case_insensitive_by_default(self) -> None:
|
||||
"""Test that matching is case-insensitive by default."""
|
||||
validator = ExcludesValidator(patterns=["BAD"])
|
||||
context = ValidationContext()
|
||||
result = validator.check("This is bad text.", context)
|
||||
|
||||
assert result.passed is False
|
||||
|
||||
def test_case_sensitive(self) -> None:
|
||||
def test_excludes_validator_case_sensitive(self) -> None:
|
||||
"""Test case-sensitive matching."""
|
||||
validator = ExcludesValidator(patterns=["BAD"], case_sensitive=True)
|
||||
context = ValidationContext()
|
||||
result = validator.check("This is bad text.", context)
|
||||
|
||||
assert result.passed is True
|
||||
|
||||
def test_regex_patterns(self) -> None:
|
||||
def test_excludes_validator_regex_patterns(self) -> None:
|
||||
"""Test regex pattern matching."""
|
||||
validator = ExcludesValidator(patterns=[r"\b\d{4}\b"]) # 4-digit numbers
|
||||
context = ValidationContext()
|
||||
|
||||
@@ -286,15 +327,18 @@ class TestExcludesValidator:
|
||||
result = validator.check("No numbers here", context)
|
||||
assert result.passed is True
|
||||
|
||||
def test_rejects_empty_patterns(self) -> None:
|
||||
def test_excludes_validator_raises_on_empty_patterns(self) -> None:
|
||||
"""Test that empty patterns list raises error."""
|
||||
with pytest.raises(InvalidThresholdError, match="cannot be empty"):
|
||||
ExcludesValidator(patterns=[])
|
||||
|
||||
def test_rejects_invalid_regex(self) -> None:
|
||||
def test_excludes_validator_raises_on_invalid_regex(self) -> None:
|
||||
"""Test that invalid regex pattern raises error at init time."""
|
||||
with pytest.raises(InvalidThresholdError, match="Invalid regex"):
|
||||
ExcludesValidator(patterns=[r"[invalid"])
|
||||
|
||||
def test_excludes_factory(self) -> None:
|
||||
def test_excludes_factory_function(self) -> None:
|
||||
"""Test the excludes() factory function."""
|
||||
validator = excludes(patterns=["test"], case_sensitive=True)
|
||||
assert isinstance(validator, ExcludesValidator)
|
||||
assert validator.name == "excludes"
|
||||
|
||||
@@ -9,7 +9,10 @@ from veritext.validators.metric import BleuValidator, LexicalValidator, RougeVal
|
||||
|
||||
|
||||
class TestBleuValidator:
|
||||
def test_bleu_passes_above_threshold(self) -> None:
|
||||
"""Tests for BleuValidator."""
|
||||
|
||||
def test_bleu_validator_passes_when_score_meets_threshold(self) -> None:
|
||||
"""Test that validator passes when BLEU score meets threshold."""
|
||||
validator = BleuValidator(min_score=0.5, variant=4)
|
||||
context = ValidationContext(reference="the cat sat on the mat")
|
||||
result = validator.check("the cat sat on the mat", context)
|
||||
@@ -19,7 +22,8 @@ class TestBleuValidator:
|
||||
assert result.actual == 1.0 # Identical text
|
||||
assert result.threshold == 0.5
|
||||
|
||||
def test_bleu_fails_below_threshold(self) -> None:
|
||||
def test_bleu_validator_fails_when_score_below_threshold(self) -> None:
|
||||
"""Test that validator fails when BLEU score is below threshold."""
|
||||
validator = BleuValidator(min_score=0.9, variant=4)
|
||||
context = ValidationContext(reference="the cat sat on the mat")
|
||||
result = validator.check("a dog ran through the park", context)
|
||||
@@ -29,7 +33,8 @@ class TestBleuValidator:
|
||||
assert result.actual < 0.9
|
||||
assert "below minimum" in result.message
|
||||
|
||||
def test_bleu_variant_selection(self) -> None:
|
||||
def test_bleu_validator_variant_selection(self) -> None:
|
||||
"""Test different BLEU variants."""
|
||||
context = ValidationContext(reference="the quick brown fox jumps")
|
||||
|
||||
for variant in (1, 2, 3, 4):
|
||||
@@ -37,32 +42,39 @@ class TestBleuValidator:
|
||||
result = validator.check("the quick brown fox", context)
|
||||
assert result.name == f"bleu-{variant}"
|
||||
|
||||
def test_bleu_requires_reference(self) -> None:
|
||||
def test_bleu_validator_raises_on_missing_reference(self) -> None:
|
||||
"""Test that validator raises when reference is missing."""
|
||||
validator = BleuValidator(min_score=0.5)
|
||||
context = ValidationContext()
|
||||
|
||||
with pytest.raises(ValidationError, match="requires reference text"):
|
||||
validator.check("some text", context)
|
||||
|
||||
def test_bleu_rejects_invalid_score(self) -> None:
|
||||
def test_bleu_validator_raises_on_invalid_min_score(self) -> None:
|
||||
"""Test that invalid min_score raises error."""
|
||||
with pytest.raises(InvalidThresholdError, match=r"between 0\.0 and 1\.0"):
|
||||
BleuValidator(min_score=1.5)
|
||||
|
||||
with pytest.raises(InvalidThresholdError, match=r"between 0\.0 and 1\.0"):
|
||||
BleuValidator(min_score=-0.1)
|
||||
|
||||
def test_bleu_rejects_invalid_variant(self) -> None:
|
||||
def test_bleu_validator_raises_on_invalid_variant(self) -> None:
|
||||
"""Test that invalid variant raises error."""
|
||||
with pytest.raises(InvalidThresholdError, match="variant must be"):
|
||||
BleuValidator(min_score=0.5, variant=5) # type: ignore[arg-type]
|
||||
|
||||
def test_bleu_factory(self) -> None:
|
||||
def test_bleu_factory_function(self) -> None:
|
||||
"""Test the bleu() factory function."""
|
||||
validator = bleu(min_score=0.6, variant=2)
|
||||
assert isinstance(validator, BleuValidator)
|
||||
assert validator.name == "bleu-2"
|
||||
|
||||
|
||||
class TestRougeValidator:
|
||||
def test_rouge_passes_above_threshold(self) -> None:
|
||||
"""Tests for RougeValidator."""
|
||||
|
||||
def test_rouge_validator_passes_when_score_meets_threshold(self) -> None:
|
||||
"""Test that validator passes when ROUGE score meets threshold."""
|
||||
validator = RougeValidator(min_score=0.5, variant="l")
|
||||
context = ValidationContext(reference="the cat sat on the mat")
|
||||
result = validator.check("the cat sat on the mat", context)
|
||||
@@ -72,7 +84,8 @@ class TestRougeValidator:
|
||||
assert result.actual == 1.0 # Identical text
|
||||
assert result.threshold == 0.5
|
||||
|
||||
def test_rouge_fails_below_threshold(self) -> None:
|
||||
def test_rouge_validator_fails_when_score_below_threshold(self) -> None:
|
||||
"""Test that validator fails when ROUGE score is below threshold."""
|
||||
validator = RougeValidator(min_score=0.9, variant="l")
|
||||
context = ValidationContext(reference="the cat sat on the mat")
|
||||
result = validator.check("a dog ran through the park", context)
|
||||
@@ -81,7 +94,8 @@ class TestRougeValidator:
|
||||
assert result.actual < 0.9
|
||||
assert "below minimum" in result.message
|
||||
|
||||
def test_rouge_variant_selection(self) -> None:
|
||||
def test_rouge_validator_variant_selection(self) -> None:
|
||||
"""Test different ROUGE variants."""
|
||||
context = ValidationContext(reference="the quick brown fox jumps")
|
||||
|
||||
for variant in ("1", "2", "l"):
|
||||
@@ -89,29 +103,36 @@ class TestRougeValidator:
|
||||
result = validator.check("the quick brown fox", context)
|
||||
assert result.name == f"rouge-{variant}"
|
||||
|
||||
def test_rouge_requires_reference(self) -> None:
|
||||
def test_rouge_validator_raises_on_missing_reference(self) -> None:
|
||||
"""Test that validator raises when reference is missing."""
|
||||
validator = RougeValidator(min_score=0.5)
|
||||
context = ValidationContext()
|
||||
|
||||
with pytest.raises(ValidationError, match="requires reference text"):
|
||||
validator.check("some text", context)
|
||||
|
||||
def test_rouge_rejects_invalid_score(self) -> None:
|
||||
def test_rouge_validator_raises_on_invalid_min_score(self) -> None:
|
||||
"""Test that invalid min_score raises error."""
|
||||
with pytest.raises(InvalidThresholdError, match=r"between 0\.0 and 1\.0"):
|
||||
RougeValidator(min_score=1.5)
|
||||
|
||||
def test_rouge_rejects_invalid_variant(self) -> None:
|
||||
def test_rouge_validator_raises_on_invalid_variant(self) -> None:
|
||||
"""Test that invalid variant raises error."""
|
||||
with pytest.raises(InvalidThresholdError, match="variant must be"):
|
||||
RougeValidator(min_score=0.5, variant="3") # type: ignore[arg-type]
|
||||
|
||||
def test_rouge_factory(self) -> None:
|
||||
def test_rouge_factory_function(self) -> None:
|
||||
"""Test the rouge() factory function."""
|
||||
validator = rouge(min_score=0.6, variant="2")
|
||||
assert isinstance(validator, RougeValidator)
|
||||
assert validator.name == "rouge-2"
|
||||
|
||||
|
||||
class TestLexicalValidator:
|
||||
def test_lexical_passes_jaccard(self) -> None:
|
||||
"""Tests for LexicalValidator."""
|
||||
|
||||
def test_lexical_validator_passes_on_jaccard(self) -> None:
|
||||
"""Test that validator passes when Jaccard similarity meets threshold."""
|
||||
validator = LexicalValidator(min_jaccard=0.5)
|
||||
context = ValidationContext(reference="the cat sat on the mat")
|
||||
result = validator.check("the cat sat on the mat", context)
|
||||
@@ -120,7 +141,8 @@ class TestLexicalValidator:
|
||||
assert result.name == "lexical"
|
||||
assert result.actual["jaccard"] == 1.0
|
||||
|
||||
def test_lexical_fails_jaccard(self) -> None:
|
||||
def test_lexical_validator_fails_on_jaccard(self) -> None:
|
||||
"""Test that validator fails when Jaccard is below threshold."""
|
||||
validator = LexicalValidator(min_jaccard=0.9)
|
||||
context = ValidationContext(reference="the cat sat on the mat")
|
||||
result = validator.check("a dog ran through the park", context)
|
||||
@@ -129,7 +151,8 @@ class TestLexicalValidator:
|
||||
assert "Jaccard" in result.message
|
||||
assert "below minimum" in result.message
|
||||
|
||||
def test_lexical_passes_overlap(self) -> None:
|
||||
def test_lexical_validator_passes_on_overlap(self) -> None:
|
||||
"""Test that validator passes when token overlap meets threshold."""
|
||||
validator = LexicalValidator(min_overlap=0.5)
|
||||
context = ValidationContext(reference="the cat sat on the mat")
|
||||
result = validator.check("the cat sat on the mat", context)
|
||||
@@ -137,7 +160,8 @@ class TestLexicalValidator:
|
||||
assert result.passed is True
|
||||
assert result.actual["token_overlap"] == 1.0
|
||||
|
||||
def test_lexical_fails_overlap(self) -> None:
|
||||
def test_lexical_validator_fails_on_overlap(self) -> None:
|
||||
"""Test that validator fails when overlap is below threshold."""
|
||||
validator = LexicalValidator(min_overlap=0.9)
|
||||
context = ValidationContext(reference="the cat sat on the mat")
|
||||
result = validator.check("a dog ran through", context)
|
||||
@@ -145,7 +169,8 @@ class TestLexicalValidator:
|
||||
assert result.passed is False
|
||||
assert "overlap" in result.message
|
||||
|
||||
def test_lexical_both_thresholds(self) -> None:
|
||||
def test_lexical_validator_with_both_thresholds(self) -> None:
|
||||
"""Test validator with both Jaccard and overlap thresholds."""
|
||||
validator = LexicalValidator(min_jaccard=0.3, min_overlap=0.5)
|
||||
context = ValidationContext(reference="the cat sat on the mat")
|
||||
result = validator.check("the cat sat", context)
|
||||
@@ -154,26 +179,31 @@ class TestLexicalValidator:
|
||||
assert "min_jaccard" in result.threshold
|
||||
assert "min_overlap" in result.threshold
|
||||
|
||||
def test_lexical_needs_threshold(self) -> None:
|
||||
def test_lexical_validator_raises_when_no_threshold(self) -> None:
|
||||
"""Test that validator raises when no threshold is provided."""
|
||||
with pytest.raises(InvalidThresholdError, match="At least one"):
|
||||
LexicalValidator()
|
||||
|
||||
def test_lexical_rejects_bad_jaccard(self) -> None:
|
||||
def test_lexical_validator_raises_on_invalid_jaccard(self) -> None:
|
||||
"""Test that invalid Jaccard threshold raises error."""
|
||||
with pytest.raises(InvalidThresholdError, match="min_jaccard"):
|
||||
LexicalValidator(min_jaccard=1.5)
|
||||
|
||||
def test_lexical_rejects_bad_overlap(self) -> None:
|
||||
def test_lexical_validator_raises_on_invalid_overlap(self) -> None:
|
||||
"""Test that invalid overlap threshold raises error."""
|
||||
with pytest.raises(InvalidThresholdError, match="min_overlap"):
|
||||
LexicalValidator(min_overlap=-0.1)
|
||||
|
||||
def test_lexical_requires_reference(self) -> None:
|
||||
def test_lexical_validator_raises_on_missing_reference(self) -> None:
|
||||
"""Test that validator raises when reference is missing."""
|
||||
validator = LexicalValidator(min_jaccard=0.5)
|
||||
context = ValidationContext()
|
||||
|
||||
with pytest.raises(ValidationError, match="requires reference text"):
|
||||
validator.check("some text", context)
|
||||
|
||||
def test_lexical_factory(self) -> None:
|
||||
def test_lexical_factory_function(self) -> None:
|
||||
"""Test the lexical() factory function."""
|
||||
validator = lexical(min_jaccard=0.5, min_overlap=0.6)
|
||||
assert isinstance(validator, LexicalValidator)
|
||||
assert validator.name == "lexical"
|
||||
@@ -181,11 +211,15 @@ class TestLexicalValidator:
|
||||
|
||||
# SemanticValidator tests - conditionally run if sentence-transformers is installed
|
||||
class TestSemanticValidator:
|
||||
"""Tests for SemanticValidator."""
|
||||
|
||||
@staticmethod
|
||||
def _skip_if_no_transformers() -> None:
|
||||
"""Skip test if sentence-transformers is not installed."""
|
||||
pytest.importorskip("sentence_transformers")
|
||||
|
||||
def test_semantic_passes_above_threshold(self) -> None:
|
||||
def test_semantic_validator_passes_when_score_meets_threshold(self) -> None:
|
||||
"""Test that validator passes when semantic similarity meets threshold."""
|
||||
self._skip_if_no_transformers()
|
||||
from veritext.validators.metric import SemanticValidator
|
||||
|
||||
@@ -198,7 +232,8 @@ class TestSemanticValidator:
|
||||
assert result.actual >= 0.99 # Identical text
|
||||
assert result.threshold == 0.5
|
||||
|
||||
def test_semantic_fails_below_threshold(self) -> None:
|
||||
def test_semantic_validator_fails_when_score_below_threshold(self) -> None:
|
||||
"""Test that validator fails when semantic similarity is below threshold."""
|
||||
self._skip_if_no_transformers()
|
||||
from veritext.validators.metric import SemanticValidator
|
||||
|
||||
@@ -213,7 +248,8 @@ class TestSemanticValidator:
|
||||
assert result.actual < 0.99
|
||||
assert "below minimum" in result.message
|
||||
|
||||
def test_semantic_requires_reference(self) -> None:
|
||||
def test_semantic_validator_raises_on_missing_reference(self) -> None:
|
||||
"""Test that validator raises when reference is missing."""
|
||||
self._skip_if_no_transformers()
|
||||
from veritext.validators.metric import SemanticValidator
|
||||
|
||||
@@ -223,7 +259,9 @@ class TestSemanticValidator:
|
||||
with pytest.raises(ValidationError, match="requires reference text"):
|
||||
validator.check("some text", context)
|
||||
|
||||
def test_semantic_rejects_invalid_score(self) -> None:
|
||||
def test_semantic_validator_raises_on_invalid_min_score(self) -> None:
|
||||
"""Test that invalid min_score raises error without loading model."""
|
||||
# This test doesn't need sentence-transformers since validation happens first
|
||||
with pytest.raises(InvalidThresholdError, match=r"between 0\.0 and 1\.0"):
|
||||
from veritext.validators.metric import SemanticValidator
|
||||
|
||||
@@ -234,7 +272,8 @@ class TestSemanticValidator:
|
||||
|
||||
SemanticValidator(min_score=-0.1)
|
||||
|
||||
def test_semantic_factory(self) -> None:
|
||||
def test_semantic_factory_function(self) -> None:
|
||||
"""Test the semantic() factory function."""
|
||||
self._skip_if_no_transformers()
|
||||
from veritext.validators import semantic
|
||||
from veritext.validators.metric import SemanticValidator
|
||||
|
||||
Reference in New Issue
Block a user