evalwise 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,854 @@
1
+ """
2
+ Text assertions - deterministic checks for LLM text outputs.
3
+
4
+ These assertions are designed to be:
5
+ 1. Deterministic - same input always gives same result
6
+ 2. Fast - no LLM calls required
7
+ 3. Interpretable - clear pass/fail with explanation
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import json
13
+ import re
14
+ import time
15
+ from typing import Any
16
+
17
+ from evalwise.core.context import get_context
18
+ from evalwise.core.result import AssertionResult, Status
19
+
20
+
21
+ class Assert:
22
+ """
23
+ Deterministic assertions for text outputs.
24
+
25
+ Example:
26
+ from evalwise import Assert
27
+
28
+ # Basic checks
29
+ Assert.contains(response, "hello")
30
+ Assert.word_count(response, min=10, max=100)
31
+ Assert.json_valid(response)
32
+
33
+ # Factuality
34
+ Assert.entails(response, source=document)
35
+ Assert.no_hallucinated_urls(response)
36
+ """
37
+
38
+ # ==================== Format & Structure ====================
39
+
40
+ @staticmethod
41
+ def contains(text: str, substring: str, case_sensitive: bool = True) -> bool:
42
+ """Assert that text contains a substring."""
43
+ start = time.perf_counter()
44
+
45
+ if case_sensitive:
46
+ passed = substring in text
47
+ else:
48
+ passed = substring.lower() in text.lower()
49
+
50
+ result = AssertionResult(
51
+ name="contains",
52
+ status=Status.PASS if passed else Status.FAIL,
53
+ message=f"Expected to contain: {substring!r}" if not passed else None,
54
+ expected=substring,
55
+ actual=text[:200] + "..." if len(text) > 200 else text,
56
+ duration_ms=(time.perf_counter() - start) * 1000,
57
+ )
58
+ get_context().add_assertion(result)
59
+ return passed
60
+
61
+ @staticmethod
62
+ def not_contains(text: str, substring: str, case_sensitive: bool = True) -> bool:
63
+ """Assert that text does NOT contain a substring."""
64
+ start = time.perf_counter()
65
+
66
+ if case_sensitive:
67
+ passed = substring not in text
68
+ else:
69
+ passed = substring.lower() not in text.lower()
70
+
71
+ result = AssertionResult(
72
+ name="not_contains",
73
+ status=Status.PASS if passed else Status.FAIL,
74
+ message=f"Should not contain: {substring!r}" if not passed else None,
75
+ expected=f"not {substring!r}",
76
+ actual=text[:200] + "..." if len(text) > 200 else text,
77
+ duration_ms=(time.perf_counter() - start) * 1000,
78
+ )
79
+ get_context().add_assertion(result)
80
+ return passed
81
+
82
+ @staticmethod
83
+ def equals(text: str, expected: str, strip: bool = True) -> bool:
84
+ """Assert exact string equality."""
85
+ start = time.perf_counter()
86
+
87
+ actual = text.strip() if strip else text
88
+ exp = expected.strip() if strip else expected
89
+ passed = actual == exp
90
+
91
+ result = AssertionResult(
92
+ name="equals",
93
+ status=Status.PASS if passed else Status.FAIL,
94
+ message=f"Expected exact match" if not passed else None,
95
+ expected=expected,
96
+ actual=text,
97
+ duration_ms=(time.perf_counter() - start) * 1000,
98
+ )
99
+ get_context().add_assertion(result)
100
+ return passed
101
+
102
+ @staticmethod
103
+ def regex(text: str, pattern: str, flags: int = 0) -> bool:
104
+ """Assert that text matches a regex pattern."""
105
+ start = time.perf_counter()
106
+
107
+ match = re.search(pattern, text, flags)
108
+ passed = match is not None
109
+
110
+ result = AssertionResult(
111
+ name="regex",
112
+ status=Status.PASS if passed else Status.FAIL,
113
+ message=f"Pattern not found: {pattern}" if not passed else None,
114
+ expected=pattern,
115
+ actual=match.group() if match else None,
116
+ duration_ms=(time.perf_counter() - start) * 1000,
117
+ )
118
+ get_context().add_assertion(result)
119
+ return passed
120
+
121
+ @staticmethod
122
+ def starts_with(text: str, prefix: str) -> bool:
123
+ """Assert that text starts with prefix."""
124
+ start = time.perf_counter()
125
+ passed = text.startswith(prefix)
126
+
127
+ result = AssertionResult(
128
+ name="starts_with",
129
+ status=Status.PASS if passed else Status.FAIL,
130
+ message=f"Expected to start with: {prefix!r}" if not passed else None,
131
+ expected=prefix,
132
+ actual=text[:len(prefix) + 20],
133
+ duration_ms=(time.perf_counter() - start) * 1000,
134
+ )
135
+ get_context().add_assertion(result)
136
+ return passed
137
+
138
+ @staticmethod
139
+ def ends_with(text: str, suffix: str) -> bool:
140
+ """Assert that text ends with suffix."""
141
+ start = time.perf_counter()
142
+ passed = text.endswith(suffix)
143
+
144
+ result = AssertionResult(
145
+ name="ends_with",
146
+ status=Status.PASS if passed else Status.FAIL,
147
+ message=f"Expected to end with: {suffix!r}" if not passed else None,
148
+ expected=suffix,
149
+ actual=text[-len(suffix) - 20:],
150
+ duration_ms=(time.perf_counter() - start) * 1000,
151
+ )
152
+ get_context().add_assertion(result)
153
+ return passed
154
+
155
+ # ==================== JSON & Structured ====================
156
+
157
+ @staticmethod
158
+ def json_valid(text: str) -> bool:
159
+ """Assert that text is valid JSON."""
160
+ start = time.perf_counter()
161
+
162
+ try:
163
+ json.loads(text)
164
+ passed = True
165
+ message = None
166
+ except json.JSONDecodeError as e:
167
+ passed = False
168
+ message = f"Invalid JSON: {e}"
169
+
170
+ result = AssertionResult(
171
+ name="json_valid",
172
+ status=Status.PASS if passed else Status.FAIL,
173
+ message=message,
174
+ duration_ms=(time.perf_counter() - start) * 1000,
175
+ )
176
+ get_context().add_assertion(result)
177
+ return passed
178
+
179
+ @staticmethod
180
+ def json_schema(text: str, schema: dict[str, Any]) -> bool:
181
+ """Assert that JSON text matches a JSON schema."""
182
+ start = time.perf_counter()
183
+
184
+ try:
185
+ from jsonschema import validate, ValidationError
186
+ except ImportError:
187
+ raise ImportError("jsonschema required: pip install jsonschema")
188
+
189
+ try:
190
+ data = json.loads(text)
191
+ validate(instance=data, schema=schema)
192
+ passed = True
193
+ message = None
194
+ except json.JSONDecodeError as e:
195
+ passed = False
196
+ message = f"Invalid JSON: {e}"
197
+ except ValidationError as e:
198
+ passed = False
199
+ message = f"Schema validation failed: {e.message}"
200
+
201
+ result = AssertionResult(
202
+ name="json_schema",
203
+ status=Status.PASS if passed else Status.FAIL,
204
+ message=message,
205
+ expected=schema,
206
+ duration_ms=(time.perf_counter() - start) * 1000,
207
+ )
208
+ get_context().add_assertion(result)
209
+ return passed
210
+
211
+ @staticmethod
212
+ def json_has_keys(text: str, keys: list[str]) -> bool:
213
+ """Assert that JSON object has required keys."""
214
+ start = time.perf_counter()
215
+
216
+ try:
217
+ data = json.loads(text)
218
+ if not isinstance(data, dict):
219
+ passed = False
220
+ message = "JSON is not an object"
221
+ missing = keys
222
+ else:
223
+ missing = [k for k in keys if k not in data]
224
+ passed = len(missing) == 0
225
+ message = f"Missing keys: {missing}" if missing else None
226
+ except json.JSONDecodeError as e:
227
+ passed = False
228
+ message = f"Invalid JSON: {e}"
229
+ missing = keys
230
+
231
+ result = AssertionResult(
232
+ name="json_has_keys",
233
+ status=Status.PASS if passed else Status.FAIL,
234
+ message=message,
235
+ expected=keys,
236
+ actual=missing if not passed else keys,
237
+ duration_ms=(time.perf_counter() - start) * 1000,
238
+ )
239
+ get_context().add_assertion(result)
240
+ return passed
241
+
242
+ # ==================== Length & Count ====================
243
+
244
+ @staticmethod
245
+ def word_count(
246
+ text: str,
247
+ *,
248
+ min: int | None = None,
249
+ max: int | None = None,
250
+ exactly: int | None = None,
251
+ ) -> bool:
252
+ """Assert word count is within bounds."""
253
+ start = time.perf_counter()
254
+
255
+ words = len(text.split())
256
+
257
+ if exactly is not None:
258
+ passed = words == exactly
259
+ expected = f"exactly {exactly}"
260
+ elif min is not None and max is not None:
261
+ passed = min <= words <= max
262
+ expected = f"{min}-{max}"
263
+ elif min is not None:
264
+ passed = words >= min
265
+ expected = f">= {min}"
266
+ elif max is not None:
267
+ passed = words <= max
268
+ expected = f"<= {max}"
269
+ else:
270
+ passed = True
271
+ expected = "any"
272
+
273
+ result = AssertionResult(
274
+ name="word_count",
275
+ status=Status.PASS if passed else Status.FAIL,
276
+ message=f"Word count {words} not in range {expected}" if not passed else None,
277
+ expected=expected,
278
+ actual=words,
279
+ duration_ms=(time.perf_counter() - start) * 1000,
280
+ )
281
+ get_context().add_assertion(result)
282
+ return passed
283
+
284
+ @staticmethod
285
+ def char_count(
286
+ text: str,
287
+ *,
288
+ min: int | None = None,
289
+ max: int | None = None,
290
+ ) -> bool:
291
+ """Assert character count is within bounds."""
292
+ start = time.perf_counter()
293
+
294
+ chars = len(text)
295
+
296
+ if min is not None and max is not None:
297
+ passed = min <= chars <= max
298
+ expected = f"{min}-{max}"
299
+ elif min is not None:
300
+ passed = chars >= min
301
+ expected = f">= {min}"
302
+ elif max is not None:
303
+ passed = chars <= max
304
+ expected = f"<= {max}"
305
+ else:
306
+ passed = True
307
+ expected = "any"
308
+
309
+ result = AssertionResult(
310
+ name="char_count",
311
+ status=Status.PASS if passed else Status.FAIL,
312
+ message=f"Char count {chars} not in range {expected}" if not passed else None,
313
+ expected=expected,
314
+ actual=chars,
315
+ duration_ms=(time.perf_counter() - start) * 1000,
316
+ )
317
+ get_context().add_assertion(result)
318
+ return passed
319
+
320
+ @staticmethod
321
+ def sentence_count(
322
+ text: str,
323
+ *,
324
+ min: int | None = None,
325
+ max: int | None = None,
326
+ exactly: int | None = None,
327
+ ) -> bool:
328
+ """Assert sentence count is within bounds."""
329
+ start = time.perf_counter()
330
+
331
+ # Simple sentence splitting on .!?
332
+ sentences = len(re.findall(r'[.!?]+', text))
333
+
334
+ if exactly is not None:
335
+ passed = sentences == exactly
336
+ expected = f"exactly {exactly}"
337
+ elif min is not None and max is not None:
338
+ passed = min <= sentences <= max
339
+ expected = f"{min}-{max}"
340
+ elif min is not None:
341
+ passed = sentences >= min
342
+ expected = f">= {min}"
343
+ elif max is not None:
344
+ passed = sentences <= max
345
+ expected = f"<= {max}"
346
+ else:
347
+ passed = True
348
+ expected = "any"
349
+
350
+ result = AssertionResult(
351
+ name="sentence_count",
352
+ status=Status.PASS if passed else Status.FAIL,
353
+ message=f"Sentence count {sentences} not in range {expected}" if not passed else None,
354
+ expected=expected,
355
+ actual=sentences,
356
+ duration_ms=(time.perf_counter() - start) * 1000,
357
+ )
358
+ get_context().add_assertion(result)
359
+ return passed
360
+
361
+ @staticmethod
362
+ def line_count(
363
+ text: str,
364
+ *,
365
+ min: int | None = None,
366
+ max: int | None = None,
367
+ exactly: int | None = None,
368
+ ) -> bool:
369
+ """Assert line count is within bounds."""
370
+ start = time.perf_counter()
371
+
372
+ lines = len(text.strip().split('\n'))
373
+
374
+ if exactly is not None:
375
+ passed = lines == exactly
376
+ expected = f"exactly {exactly}"
377
+ elif min is not None and max is not None:
378
+ passed = min <= lines <= max
379
+ expected = f"{min}-{max}"
380
+ elif min is not None:
381
+ passed = lines >= min
382
+ expected = f">= {min}"
383
+ elif max is not None:
384
+ passed = lines <= max
385
+ expected = f"<= {max}"
386
+ else:
387
+ passed = True
388
+ expected = "any"
389
+
390
+ result = AssertionResult(
391
+ name="line_count",
392
+ status=Status.PASS if passed else Status.FAIL,
393
+ message=f"Line count {lines} not in range {expected}" if not passed else None,
394
+ expected=expected,
395
+ actual=lines,
396
+ duration_ms=(time.perf_counter() - start) * 1000,
397
+ )
398
+ get_context().add_assertion(result)
399
+ return passed
400
+
401
+ @staticmethod
402
+ def bullet_count(
403
+ text: str,
404
+ *,
405
+ min: int | None = None,
406
+ max: int | None = None,
407
+ exactly: int | None = None,
408
+ ) -> bool:
409
+ """Assert bullet point count (lines starting with -, *, •, or numbers)."""
410
+ start = time.perf_counter()
411
+
412
+ bullet_pattern = r'^[\s]*[-*•]|\d+[.\)]'
413
+ bullets = len(re.findall(bullet_pattern, text, re.MULTILINE))
414
+
415
+ if exactly is not None:
416
+ passed = bullets == exactly
417
+ expected = f"exactly {exactly}"
418
+ elif min is not None and max is not None:
419
+ passed = min <= bullets <= max
420
+ expected = f"{min}-{max}"
421
+ elif min is not None:
422
+ passed = bullets >= min
423
+ expected = f">= {min}"
424
+ elif max is not None:
425
+ passed = bullets <= max
426
+ expected = f"<= {max}"
427
+ else:
428
+ passed = True
429
+ expected = "any"
430
+
431
+ result = AssertionResult(
432
+ name="bullet_count",
433
+ status=Status.PASS if passed else Status.FAIL,
434
+ message=f"Bullet count {bullets} not in range {expected}" if not passed else None,
435
+ expected=expected,
436
+ actual=bullets,
437
+ duration_ms=(time.perf_counter() - start) * 1000,
438
+ )
439
+ get_context().add_assertion(result)
440
+ return passed
441
+
442
+ # ==================== Readability ====================
443
+
444
+ @staticmethod
445
+ def readability(text: str, *, min_score: float | None = None, max_score: float | None = None) -> bool:
446
+ """
447
+ Assert Flesch Reading Ease score is within bounds.
448
+
449
+ Scores: 0-30 (very difficult) to 90-100 (very easy)
450
+ - 60-70: Standard / Plain English
451
+ - 70-80: Fairly easy
452
+ - 80-90: Easy
453
+ """
454
+ start = time.perf_counter()
455
+
456
+ try:
457
+ import textstat
458
+ except ImportError:
459
+ raise ImportError("textstat required: pip install textstat")
460
+
461
+ score = textstat.flesch_reading_ease(text)
462
+
463
+ if min_score is not None and max_score is not None:
464
+ passed = min_score <= score <= max_score
465
+ expected = f"{min_score}-{max_score}"
466
+ elif min_score is not None:
467
+ passed = score >= min_score
468
+ expected = f">= {min_score}"
469
+ elif max_score is not None:
470
+ passed = score <= max_score
471
+ expected = f"<= {max_score}"
472
+ else:
473
+ passed = True
474
+ expected = "any"
475
+
476
+ result = AssertionResult(
477
+ name="readability",
478
+ status=Status.PASS if passed else Status.FAIL,
479
+ message=f"Flesch score {score:.1f} not in range {expected}" if not passed else None,
480
+ expected=expected,
481
+ actual=round(score, 1),
482
+ duration_ms=(time.perf_counter() - start) * 1000,
483
+ )
484
+ get_context().add_assertion(result)
485
+ return passed
486
+
487
+ @staticmethod
488
+ def grade_level(text: str, *, min_grade: float | None = None, max_grade: float | None = None) -> bool:
489
+ """
490
+ Assert Flesch-Kincaid Grade Level is within bounds.
491
+
492
+ Score represents US school grade level needed to understand the text.
493
+ """
494
+ start = time.perf_counter()
495
+
496
+ try:
497
+ import textstat
498
+ except ImportError:
499
+ raise ImportError("textstat required: pip install textstat")
500
+
501
+ grade = textstat.flesch_kincaid_grade(text)
502
+
503
+ if min_grade is not None and max_grade is not None:
504
+ passed = min_grade <= grade <= max_grade
505
+ expected = f"grade {min_grade}-{max_grade}"
506
+ elif min_grade is not None:
507
+ passed = grade >= min_grade
508
+ expected = f">= grade {min_grade}"
509
+ elif max_grade is not None:
510
+ passed = grade <= max_grade
511
+ expected = f"<= grade {max_grade}"
512
+ else:
513
+ passed = True
514
+ expected = "any"
515
+
516
+ result = AssertionResult(
517
+ name="grade_level",
518
+ status=Status.PASS if passed else Status.FAIL,
519
+ message=f"Grade level {grade:.1f} not in range {expected}" if not passed else None,
520
+ expected=expected,
521
+ actual=round(grade, 1),
522
+ duration_ms=(time.perf_counter() - start) * 1000,
523
+ )
524
+ get_context().add_assertion(result)
525
+ return passed
526
+
527
+ # ==================== Language ====================
528
+
529
+ @staticmethod
530
+ def language_is(text: str, expected_lang: str) -> bool:
531
+ """Assert text is in expected language (ISO 639-1 code like 'en', 'es', 'fr')."""
532
+ start = time.perf_counter()
533
+
534
+ try:
535
+ from langdetect import detect, LangDetectException
536
+ except ImportError:
537
+ raise ImportError("langdetect required: pip install langdetect")
538
+
539
+ try:
540
+ detected = detect(text)
541
+ passed = detected == expected_lang
542
+ message = f"Detected '{detected}', expected '{expected_lang}'" if not passed else None
543
+ except LangDetectException as e:
544
+ passed = False
545
+ detected = None
546
+ message = f"Language detection failed: {e}"
547
+
548
+ result = AssertionResult(
549
+ name="language_is",
550
+ status=Status.PASS if passed else Status.FAIL,
551
+ message=message,
552
+ expected=expected_lang,
553
+ actual=detected,
554
+ duration_ms=(time.perf_counter() - start) * 1000,
555
+ )
556
+ get_context().add_assertion(result)
557
+ return passed
558
+
559
+ # ==================== URLs & Links ====================
560
+
561
+ @staticmethod
562
+ def no_urls(text: str) -> bool:
563
+ """Assert text contains no URLs."""
564
+ start = time.perf_counter()
565
+
566
+ url_pattern = r'https?://[^\s<>"{}|\\^`\[\]]+'
567
+ urls = re.findall(url_pattern, text)
568
+ passed = len(urls) == 0
569
+
570
+ result = AssertionResult(
571
+ name="no_urls",
572
+ status=Status.PASS if passed else Status.FAIL,
573
+ message=f"Found URLs: {urls[:3]}" if not passed else None,
574
+ actual=urls if urls else None,
575
+ duration_ms=(time.perf_counter() - start) * 1000,
576
+ )
577
+ get_context().add_assertion(result)
578
+ return passed
579
+
580
+ @staticmethod
581
+ def urls_valid(text: str, timeout: float = 5.0) -> bool:
582
+ """Assert all URLs in text return 2xx status (no hallucinated URLs)."""
583
+ start = time.perf_counter()
584
+
585
+ try:
586
+ import httpx
587
+ except ImportError:
588
+ raise ImportError("httpx required: pip install httpx")
589
+
590
+ url_pattern = r'https?://[^\s<>"{}|\\^`\[\]]+'
591
+ urls = re.findall(url_pattern, text)
592
+
593
+ if not urls:
594
+ passed = True
595
+ invalid = []
596
+ else:
597
+ invalid = []
598
+ for url in urls[:10]: # Limit to 10 URLs
599
+ try:
600
+ resp = httpx.head(url, timeout=timeout, follow_redirects=True)
601
+ if resp.status_code >= 400:
602
+ invalid.append((url, resp.status_code))
603
+ except Exception as e:
604
+ invalid.append((url, str(e)))
605
+
606
+ passed = len(invalid) == 0
607
+
608
+ result = AssertionResult(
609
+ name="urls_valid",
610
+ status=Status.PASS if passed else Status.FAIL,
611
+ message=f"Invalid URLs: {invalid}" if invalid else None,
612
+ expected="all URLs return 2xx",
613
+ actual=invalid if invalid else f"{len(urls)} URLs valid",
614
+ duration_ms=(time.perf_counter() - start) * 1000,
615
+ )
616
+ get_context().add_assertion(result)
617
+ return passed
618
+
619
+ # ==================== Code ====================
620
+
621
+ @staticmethod
622
+ def code_parses(text: str, language: str = "python") -> bool:
623
+ """Assert code is syntactically valid."""
624
+ start = time.perf_counter()
625
+
626
+ # Extract code from markdown blocks if present
627
+ code = text
628
+ code_block = re.search(r'```(?:\w+)?\n(.*?)```', text, re.DOTALL)
629
+ if code_block:
630
+ code = code_block.group(1)
631
+
632
+ if language == "python":
633
+ try:
634
+ import ast
635
+ ast.parse(code)
636
+ passed = True
637
+ message = None
638
+ except SyntaxError as e:
639
+ passed = False
640
+ message = f"Syntax error: {e}"
641
+ elif language == "json":
642
+ try:
643
+ json.loads(code)
644
+ passed = True
645
+ message = None
646
+ except json.JSONDecodeError as e:
647
+ passed = False
648
+ message = f"JSON error: {e}"
649
+ else:
650
+ passed = True
651
+ message = f"Language '{language}' validation not implemented"
652
+
653
+ result = AssertionResult(
654
+ name="code_parses",
655
+ status=Status.PASS if passed else Status.FAIL,
656
+ message=message,
657
+ expected=f"valid {language}",
658
+ duration_ms=(time.perf_counter() - start) * 1000,
659
+ )
660
+ get_context().add_assertion(result)
661
+ return passed
662
+
663
+ @staticmethod
664
+ def code_runs(text: str, timeout: float = 5.0) -> bool:
665
+ """Assert Python code executes without error."""
666
+ start = time.perf_counter()
667
+
668
+ # Extract code from markdown blocks if present
669
+ code = text
670
+ code_block = re.search(r'```(?:python)?\n(.*?)```', text, re.DOTALL)
671
+ if code_block:
672
+ code = code_block.group(1)
673
+
674
+ try:
675
+ exec(compile(code, '<string>', 'exec'), {'__builtins__': __builtins__})
676
+ passed = True
677
+ message = None
678
+ except Exception as e:
679
+ passed = False
680
+ message = f"{type(e).__name__}: {e}"
681
+
682
+ result = AssertionResult(
683
+ name="code_runs",
684
+ status=Status.PASS if passed else Status.FAIL,
685
+ message=message,
686
+ expected="executes without error",
687
+ duration_ms=(time.perf_counter() - start) * 1000,
688
+ )
689
+ get_context().add_assertion(result)
690
+ return passed
691
+
692
+ # ==================== Semantic (using embeddings) ====================
693
+
694
+ @staticmethod
695
+ def embedding_similarity(
696
+ text: str,
697
+ reference: str,
698
+ *,
699
+ threshold: float = 0.8,
700
+ model: str = "all-MiniLM-L6-v2",
701
+ ) -> bool:
702
+ """Assert text is semantically similar to reference (cosine similarity)."""
703
+ start = time.perf_counter()
704
+
705
+ try:
706
+ from sentence_transformers import SentenceTransformer
707
+ import numpy as np
708
+ except ImportError:
709
+ raise ImportError(
710
+ "sentence-transformers required: pip install sentence-transformers"
711
+ )
712
+
713
+ # Load model (cached after first call)
714
+ encoder = SentenceTransformer(model)
715
+
716
+ # Compute embeddings
717
+ embeddings = encoder.encode([text, reference])
718
+
719
+ # Cosine similarity
720
+ similarity = float(np.dot(embeddings[0], embeddings[1]) / (
721
+ np.linalg.norm(embeddings[0]) * np.linalg.norm(embeddings[1])
722
+ ))
723
+
724
+ passed = similarity >= threshold
725
+
726
+ result = AssertionResult(
727
+ name="embedding_similarity",
728
+ status=Status.PASS if passed else Status.FAIL,
729
+ message=f"Similarity {similarity:.3f} < threshold {threshold}" if not passed else None,
730
+ expected=f">= {threshold}",
731
+ actual=round(similarity, 3),
732
+ duration_ms=(time.perf_counter() - start) * 1000,
733
+ )
734
+ get_context().add_assertion(result)
735
+ return passed
736
+
737
+ # ==================== Factuality (NLI-based) ====================
738
+
739
+ @staticmethod
740
+ def entails(text: str, *, source: str, threshold: float = 0.7) -> bool:
741
+ """
742
+ Assert text is entailed by (follows from) the source.
743
+ Uses NLI model - deterministic given the model.
744
+ """
745
+ start = time.perf_counter()
746
+
747
+ try:
748
+ import torch
749
+ from transformers import pipeline
750
+ except ImportError as e:
751
+ if "torch" in str(e) or "torch" not in dir():
752
+ raise ImportError(
753
+ "torch required for NLI: pip install torch transformers"
754
+ )
755
+ raise ImportError("transformers required: pip install transformers")
756
+
757
+ # Load NLI model (cached after first call)
758
+ nli = pipeline(
759
+ "text-classification",
760
+ model="facebook/bart-large-mnli",
761
+ device=-1, # CPU
762
+ )
763
+
764
+ # Check if source entails text
765
+ result_nli = nli(f"{source}</s></s>{text}", top_k=3)
766
+
767
+ # Find entailment score
768
+ entailment_score = 0.0
769
+ for r in result_nli:
770
+ if r["label"] == "entailment":
771
+ entailment_score = r["score"]
772
+ break
773
+
774
+ passed = entailment_score >= threshold
775
+
776
+ result = AssertionResult(
777
+ name="entails",
778
+ status=Status.PASS if passed else Status.FAIL,
779
+ message=f"Entailment score {entailment_score:.3f} < threshold {threshold}" if not passed else None,
780
+ expected=f">= {threshold}",
781
+ actual=round(entailment_score, 3),
782
+ duration_ms=(time.perf_counter() - start) * 1000,
783
+ )
784
+ get_context().add_assertion(result)
785
+ return passed
786
+
787
+ @staticmethod
788
+ def no_contradiction(text: str, *, source: str, threshold: float = 0.3) -> bool:
789
+ """
790
+ Assert text does not contradict the source.
791
+ Uses NLI model - deterministic given the model.
792
+ """
793
+ start = time.perf_counter()
794
+
795
+ try:
796
+ import torch
797
+ from transformers import pipeline
798
+ except ImportError as e:
799
+ if "torch" in str(e) or "torch" not in dir():
800
+ raise ImportError(
801
+ "torch required for NLI: pip install torch transformers"
802
+ )
803
+ raise ImportError("transformers required: pip install transformers")
804
+
805
+ nli = pipeline(
806
+ "text-classification",
807
+ model="facebook/bart-large-mnli",
808
+ device=-1,
809
+ )
810
+
811
+ result_nli = nli(f"{source}</s></s>{text}", top_k=3)
812
+
813
+ contradiction_score = 0.0
814
+ for r in result_nli:
815
+ if r["label"] == "contradiction":
816
+ contradiction_score = r["score"]
817
+ break
818
+
819
+ passed = contradiction_score < threshold
820
+
821
+ result = AssertionResult(
822
+ name="no_contradiction",
823
+ status=Status.PASS if passed else Status.FAIL,
824
+ message=f"Contradiction score {contradiction_score:.3f} >= threshold {threshold}" if not passed else None,
825
+ expected=f"< {threshold}",
826
+ actual=round(contradiction_score, 3),
827
+ duration_ms=(time.perf_counter() - start) * 1000,
828
+ )
829
+ get_context().add_assertion(result)
830
+ return passed
831
+
832
+ # ==================== Custom ====================
833
+
834
+ @staticmethod
835
+ def custom(
836
+ passed: bool,
837
+ name: str = "custom",
838
+ message: str | None = None,
839
+ expected: Any = None,
840
+ actual: Any = None,
841
+ ) -> bool:
842
+ """Register a custom assertion result."""
843
+ start = time.perf_counter()
844
+
845
+ result = AssertionResult(
846
+ name=name,
847
+ status=Status.PASS if passed else Status.FAIL,
848
+ message=message if not passed else None,
849
+ expected=expected,
850
+ actual=actual,
851
+ duration_ms=(time.perf_counter() - start) * 1000,
852
+ )
853
+ get_context().add_assertion(result)
854
+ return passed