Enhancements to LLM concurrency to improve overall throughput
This commit is contained in:
@@ -88,7 +88,8 @@ Useful configuration can be supplied by CLI flag or environment variable. CLI fl
|
|||||||
| `AUDITA_MODEL` | `--model` | `openrouter/google/gemma-4-31b-it` | LLM model name sent to the configured OpenAI-compatible endpoint |
|
| `AUDITA_MODEL` | `--model` | `openrouter/google/gemma-4-31b-it` | LLM model name sent to the configured OpenAI-compatible endpoint |
|
||||||
| `AUDITA_BASE_URL` | `--base-url` | `https://openrouter.ai/api/v1` | OpenAI-compatible API base URL |
|
| `AUDITA_BASE_URL` | `--base-url` | `https://openrouter.ai/api/v1` | OpenAI-compatible API base URL |
|
||||||
| `AUDITA_MAX_RETRIES` | `--max-retries` | `3` | Maximum Instructor retries for structured responses |
|
| `AUDITA_MAX_RETRIES` | `--max-retries` | `3` | Maximum Instructor retries for structured responses |
|
||||||
| `AUDITA_MAX_SECTION_TOKENS` | `--max-section-tokens` | `6144` | Maximum estimated tokens per transcript batch |
|
| `AUDITA_MAX_SECTION_TOKENS` | `--max-section-tokens` | `16384` | Maximum estimated tokens per proposal-stage transcript section |
|
||||||
|
| `AUDITA_MIN_SECTION_TOKENS` | `--min-section-tokens` | `4096` | Minimum estimated tokens per proposal-stage transcript section when balancing for concurrency |
|
||||||
| `AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD` | `--glossary-confidence-threshold` | `0.8` | Minimum confidence required for glossary proposals to survive validation |
|
| `AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD` | `--glossary-confidence-threshold` | `0.8` | Minimum confidence required for glossary proposals to survive validation |
|
||||||
| `AUDITA_GRAMMAR_CONFIDENCE_THRESHOLD` | `--grammar-confidence-threshold` | `0.8` | Minimum confidence required for grammar proposals to survive validation |
|
| `AUDITA_GRAMMAR_CONFIDENCE_THRESHOLD` | `--grammar-confidence-threshold` | `0.8` | Minimum confidence required for grammar proposals to survive validation |
|
||||||
| `AUDITA_HOMOPHONES_CONFIDENCE_THRESHOLD` | `--homophones-confidence-threshold` | `0.8` | Minimum confidence required for homophone proposals to survive validation |
|
| `AUDITA_HOMOPHONES_CONFIDENCE_THRESHOLD` | `--homophones-confidence-threshold` | `0.8` | Minimum confidence required for homophone proposals to survive validation |
|
||||||
|
|||||||
@@ -37,6 +37,11 @@ def _build_parser() -> argparse.ArgumentParser:
|
|||||||
process.add_argument("--base-url", help="OpenAI-compatible API base URL for Audita LLM stages")
|
process.add_argument("--base-url", help="OpenAI-compatible API base URL for Audita LLM stages")
|
||||||
process.add_argument("--max-retries", type=int, help="maximum structured-output retries for LLM stages")
|
process.add_argument("--max-retries", type=int, help="maximum structured-output retries for LLM stages")
|
||||||
process.add_argument("--max-section-tokens", type=int, help="maximum estimated tokens per transcript batch")
|
process.add_argument("--max-section-tokens", type=int, help="maximum estimated tokens per transcript batch")
|
||||||
|
process.add_argument(
|
||||||
|
"--min-section-tokens",
|
||||||
|
type=int,
|
||||||
|
help="minimum estimated tokens per transcript batch when balancing proposal-stage sections",
|
||||||
|
)
|
||||||
process.add_argument(
|
process.add_argument(
|
||||||
"--glossary-confidence-threshold",
|
"--glossary-confidence-threshold",
|
||||||
type=float,
|
type=float,
|
||||||
@@ -97,6 +102,7 @@ def _process(args: argparse.Namespace) -> int:
|
|||||||
base_url=args.base_url,
|
base_url=args.base_url,
|
||||||
max_retries=args.max_retries,
|
max_retries=args.max_retries,
|
||||||
max_section_tokens=args.max_section_tokens,
|
max_section_tokens=args.max_section_tokens,
|
||||||
|
min_section_tokens=args.min_section_tokens,
|
||||||
glossary_confidence_threshold=args.glossary_confidence_threshold,
|
glossary_confidence_threshold=args.glossary_confidence_threshold,
|
||||||
grammar_confidence_threshold=args.grammar_confidence_threshold,
|
grammar_confidence_threshold=args.grammar_confidence_threshold,
|
||||||
homophones_confidence_threshold=args.homophones_confidence_threshold,
|
homophones_confidence_threshold=args.homophones_confidence_threshold,
|
||||||
|
|||||||
@@ -117,24 +117,43 @@ class TranscriptSection:
|
|||||||
def chunk_transcript(
|
def chunk_transcript(
|
||||||
segments: List[TranscriptSegment],
|
segments: List[TranscriptSegment],
|
||||||
max_section_tokens: int,
|
max_section_tokens: int,
|
||||||
|
min_section_tokens: int = 1,
|
||||||
|
target_section_count: Optional[int] = None,
|
||||||
estimator: Optional[TokenEstimatorProtocol] = None,
|
estimator: Optional[TokenEstimatorProtocol] = None,
|
||||||
) -> List[TranscriptSection]:
|
) -> List[TranscriptSection]:
|
||||||
indexed = [IndexedSegment(index=index, segment=segment) for index, segment in enumerate(segments)]
|
indexed = [IndexedSegment(index=index, segment=segment) for index, segment in enumerate(segments)]
|
||||||
return chunk_indexed_segments(indexed, max_section_tokens, estimator=estimator)
|
return chunk_indexed_segments(
|
||||||
|
indexed,
|
||||||
|
max_section_tokens,
|
||||||
|
min_section_tokens=min_section_tokens,
|
||||||
|
target_section_count=target_section_count,
|
||||||
|
estimator=estimator,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def chunk_indexed_segments(
|
def chunk_indexed_segments(
|
||||||
indexed_segments: List[IndexedSegment],
|
indexed_segments: List[IndexedSegment],
|
||||||
max_section_tokens: int,
|
max_section_tokens: int,
|
||||||
|
min_section_tokens: int = 1,
|
||||||
|
target_section_count: Optional[int] = None,
|
||||||
estimator: Optional[TokenEstimatorProtocol] = None,
|
estimator: Optional[TokenEstimatorProtocol] = None,
|
||||||
) -> List[TranscriptSection]:
|
) -> List[TranscriptSection]:
|
||||||
batches = chunk_payload_items(
|
if target_section_count is None:
|
||||||
indexed_segments,
|
batches = chunk_payload_items(
|
||||||
max_section_tokens,
|
indexed_segments,
|
||||||
payload_fn=lambda item: item.prompt_payload(),
|
max_section_tokens,
|
||||||
estimator=estimator,
|
payload_fn=lambda item: item.prompt_payload(),
|
||||||
empty_error_message="Transcript must contain at least one segment.",
|
estimator=estimator,
|
||||||
)
|
empty_error_message="Transcript must contain at least one segment.",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
batches = _chunk_indexed_segments_for_target_count(
|
||||||
|
indexed_segments,
|
||||||
|
max_section_tokens=max_section_tokens,
|
||||||
|
min_section_tokens=min_section_tokens,
|
||||||
|
target_section_count=target_section_count,
|
||||||
|
estimator=estimator,
|
||||||
|
)
|
||||||
sections = [
|
sections = [
|
||||||
TranscriptSection(
|
TranscriptSection(
|
||||||
section_index=batch.batch_index,
|
section_index=batch.batch_index,
|
||||||
@@ -147,3 +166,137 @@ def chunk_indexed_segments(
|
|||||||
for section in sections:
|
for section in sections:
|
||||||
parse_transcript_json(section.transcript_json(), require_sequential_ids=False)
|
parse_transcript_json(section.transcript_json(), require_sequential_ids=False)
|
||||||
return sections
|
return sections
|
||||||
|
|
||||||
|
|
||||||
|
def _chunk_indexed_segments_for_target_count(
|
||||||
|
indexed_segments: List[IndexedSegment],
|
||||||
|
*,
|
||||||
|
max_section_tokens: int,
|
||||||
|
min_section_tokens: int,
|
||||||
|
target_section_count: int,
|
||||||
|
estimator: Optional[TokenEstimatorProtocol],
|
||||||
|
) -> List[TokenBatch[IndexedSegment]]:
|
||||||
|
if max_section_tokens <= 0:
|
||||||
|
raise AuditaValidationError("Maximum section token count must be greater than zero.")
|
||||||
|
if min_section_tokens <= 0:
|
||||||
|
raise AuditaValidationError("Minimum section token count must be greater than zero.")
|
||||||
|
if min_section_tokens > max_section_tokens:
|
||||||
|
raise AuditaValidationError(
|
||||||
|
"Minimum section token count must be less than or equal to maximum section token count."
|
||||||
|
)
|
||||||
|
if not indexed_segments:
|
||||||
|
raise AuditaValidationError("Transcript must contain at least one segment.")
|
||||||
|
if target_section_count <= 0:
|
||||||
|
raise AuditaValidationError("Target section count must be greater than zero.")
|
||||||
|
|
||||||
|
token_estimator = TokenEstimator() if estimator is None else estimator
|
||||||
|
single_tokens = [token_estimator.estimate_json([item.prompt_payload()]) for item in indexed_segments]
|
||||||
|
if any(tokens > max_section_tokens for tokens in single_tokens):
|
||||||
|
raise AuditaValidationError(
|
||||||
|
"A single transcript segment exceeds the maximum section token limit. "
|
||||||
|
"Raise the limit or pre-split the transcript."
|
||||||
|
)
|
||||||
|
|
||||||
|
total_tokens = sum(single_tokens)
|
||||||
|
if total_tokens < min_section_tokens:
|
||||||
|
return [_build_token_batch(indexed_segments, batch_index=0, estimator=token_estimator)]
|
||||||
|
|
||||||
|
desired_count = min(target_section_count, len(indexed_segments))
|
||||||
|
section_count = _resolve_section_count(
|
||||||
|
indexed_segments=indexed_segments,
|
||||||
|
single_tokens=single_tokens,
|
||||||
|
total_tokens=total_tokens,
|
||||||
|
desired_count=desired_count,
|
||||||
|
min_section_tokens=min_section_tokens,
|
||||||
|
max_section_tokens=max_section_tokens,
|
||||||
|
estimator=token_estimator,
|
||||||
|
)
|
||||||
|
return _build_balanced_batches(indexed_segments, single_tokens, section_count, token_estimator)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_section_count(
|
||||||
|
*,
|
||||||
|
indexed_segments: List[IndexedSegment],
|
||||||
|
single_tokens: List[int],
|
||||||
|
total_tokens: int,
|
||||||
|
desired_count: int,
|
||||||
|
min_section_tokens: int,
|
||||||
|
max_section_tokens: int,
|
||||||
|
estimator: TokenEstimatorProtocol,
|
||||||
|
) -> int:
|
||||||
|
batches = _build_balanced_batches(indexed_segments, single_tokens, desired_count, estimator)
|
||||||
|
if _batches_within_bounds(batches, min_section_tokens=min_section_tokens, max_section_tokens=max_section_tokens):
|
||||||
|
return desired_count
|
||||||
|
|
||||||
|
average_tokens = total_tokens / desired_count
|
||||||
|
if average_tokens > max_section_tokens:
|
||||||
|
candidates = range(desired_count + 1, len(indexed_segments) + 1)
|
||||||
|
elif average_tokens < min_section_tokens:
|
||||||
|
candidates = range(desired_count - 1, 0, -1)
|
||||||
|
else:
|
||||||
|
candidates = list(range(desired_count + 1, len(indexed_segments) + 1)) + list(
|
||||||
|
range(desired_count - 1, 0, -1)
|
||||||
|
)
|
||||||
|
|
||||||
|
for count in candidates:
|
||||||
|
batches = _build_balanced_batches(indexed_segments, single_tokens, count, estimator)
|
||||||
|
if _batches_within_bounds(batches, min_section_tokens=min_section_tokens, max_section_tokens=max_section_tokens):
|
||||||
|
return count
|
||||||
|
return 1
|
||||||
|
|
||||||
|
|
||||||
|
def _build_balanced_batches(
|
||||||
|
indexed_segments: List[IndexedSegment],
|
||||||
|
single_tokens: List[int],
|
||||||
|
section_count: int,
|
||||||
|
estimator: TokenEstimatorProtocol,
|
||||||
|
) -> List[TokenBatch[IndexedSegment]]:
|
||||||
|
if section_count == 1:
|
||||||
|
return [_build_token_batch(indexed_segments, batch_index=0, estimator=estimator)]
|
||||||
|
|
||||||
|
prefix_tokens = [0]
|
||||||
|
for tokens in single_tokens:
|
||||||
|
prefix_tokens.append(prefix_tokens[-1] + tokens)
|
||||||
|
|
||||||
|
cuts = [0]
|
||||||
|
total_tokens = prefix_tokens[-1]
|
||||||
|
for section_index in range(1, section_count):
|
||||||
|
target = total_tokens * section_index / section_count
|
||||||
|
min_cut = cuts[-1] + 1
|
||||||
|
max_cut = len(indexed_segments) - (section_count - section_index)
|
||||||
|
best_cut = min_cut
|
||||||
|
best_distance = None
|
||||||
|
for cut in range(min_cut, max_cut + 1):
|
||||||
|
distance = abs(prefix_tokens[cut] - target)
|
||||||
|
if best_distance is None or distance < best_distance:
|
||||||
|
best_cut = cut
|
||||||
|
best_distance = distance
|
||||||
|
cuts.append(best_cut)
|
||||||
|
cuts.append(len(indexed_segments))
|
||||||
|
|
||||||
|
return [
|
||||||
|
_build_token_batch(indexed_segments[cuts[index] : cuts[index + 1]], batch_index=index, estimator=estimator)
|
||||||
|
for index in range(section_count)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _build_token_batch(
|
||||||
|
items: Sequence[IndexedSegment],
|
||||||
|
*,
|
||||||
|
batch_index: int,
|
||||||
|
estimator: TokenEstimatorProtocol,
|
||||||
|
) -> TokenBatch[IndexedSegment]:
|
||||||
|
return TokenBatch(
|
||||||
|
batch_index=batch_index,
|
||||||
|
items=list(items),
|
||||||
|
token_count=estimator.estimate_json([item.prompt_payload() for item in items]),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _batches_within_bounds(
|
||||||
|
batches: Sequence[TokenBatch[IndexedSegment]],
|
||||||
|
*,
|
||||||
|
min_section_tokens: int,
|
||||||
|
max_section_tokens: int,
|
||||||
|
) -> bool:
|
||||||
|
return all(min_section_tokens <= batch.token_count <= max_section_tokens for batch in batches)
|
||||||
|
|||||||
@@ -13,7 +13,8 @@ DEFAULT_MODEL = "openrouter/google/gemma-4-31b-it"
|
|||||||
DEFAULT_BASE_URL = "https://openrouter.ai/api/v1"
|
DEFAULT_BASE_URL = "https://openrouter.ai/api/v1"
|
||||||
DEFAULT_LLM_CONCURRENCY = 1
|
DEFAULT_LLM_CONCURRENCY = 1
|
||||||
DEFAULT_MAX_RETRIES = 3
|
DEFAULT_MAX_RETRIES = 3
|
||||||
DEFAULT_MAX_SECTION_TOKENS = 6144
|
DEFAULT_MAX_SECTION_TOKENS = 16384
|
||||||
|
DEFAULT_MIN_SECTION_TOKENS = 4096
|
||||||
DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD = 0.80
|
DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD = 0.80
|
||||||
DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD = 0.80
|
DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD = 0.80
|
||||||
DEFAULT_HOMOPHONES_CONFIDENCE_THRESHOLD = 0.80
|
DEFAULT_HOMOPHONES_CONFIDENCE_THRESHOLD = 0.80
|
||||||
@@ -35,6 +36,7 @@ class ConfigOverrides:
|
|||||||
base_url: Optional[str] = None
|
base_url: Optional[str] = None
|
||||||
max_retries: Optional[int] = None
|
max_retries: Optional[int] = None
|
||||||
max_section_tokens: Optional[int] = None
|
max_section_tokens: Optional[int] = None
|
||||||
|
min_section_tokens: Optional[int] = None
|
||||||
glossary_confidence_threshold: Optional[float] = None
|
glossary_confidence_threshold: Optional[float] = None
|
||||||
grammar_confidence_threshold: Optional[float] = None
|
grammar_confidence_threshold: Optional[float] = None
|
||||||
homophones_confidence_threshold: Optional[float] = None
|
homophones_confidence_threshold: Optional[float] = None
|
||||||
@@ -56,6 +58,7 @@ class AuditaConfig:
|
|||||||
base_url: str = DEFAULT_BASE_URL
|
base_url: str = DEFAULT_BASE_URL
|
||||||
max_retries: int = DEFAULT_MAX_RETRIES
|
max_retries: int = DEFAULT_MAX_RETRIES
|
||||||
max_section_tokens: int = DEFAULT_MAX_SECTION_TOKENS
|
max_section_tokens: int = DEFAULT_MAX_SECTION_TOKENS
|
||||||
|
min_section_tokens: int = DEFAULT_MIN_SECTION_TOKENS
|
||||||
glossary_confidence_threshold: float = DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD
|
glossary_confidence_threshold: float = DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD
|
||||||
grammar_confidence_threshold: float = DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD
|
grammar_confidence_threshold: float = DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD
|
||||||
homophones_confidence_threshold: float = DEFAULT_HOMOPHONES_CONFIDENCE_THRESHOLD
|
homophones_confidence_threshold: float = DEFAULT_HOMOPHONES_CONFIDENCE_THRESHOLD
|
||||||
@@ -107,6 +110,12 @@ class AuditaConfig:
|
|||||||
DEFAULT_MAX_SECTION_TOKENS,
|
DEFAULT_MAX_SECTION_TOKENS,
|
||||||
"AUDITA_MAX_SECTION_TOKENS",
|
"AUDITA_MAX_SECTION_TOKENS",
|
||||||
),
|
),
|
||||||
|
min_section_tokens=_select_int(
|
||||||
|
selected.min_section_tokens,
|
||||||
|
source.get("AUDITA_MIN_SECTION_TOKENS"),
|
||||||
|
DEFAULT_MIN_SECTION_TOKENS,
|
||||||
|
"AUDITA_MIN_SECTION_TOKENS",
|
||||||
|
),
|
||||||
glossary_confidence_threshold=_select_float(
|
glossary_confidence_threshold=_select_float(
|
||||||
selected.glossary_confidence_threshold,
|
selected.glossary_confidence_threshold,
|
||||||
source.get("AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD"),
|
source.get("AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD"),
|
||||||
@@ -179,6 +188,12 @@ class AuditaConfig:
|
|||||||
raise AuditaConfigError("AUDITA_MAX_RETRIES must be greater than or equal to zero.")
|
raise AuditaConfigError("AUDITA_MAX_RETRIES must be greater than or equal to zero.")
|
||||||
if self.max_section_tokens <= 0:
|
if self.max_section_tokens <= 0:
|
||||||
raise AuditaConfigError("AUDITA_MAX_SECTION_TOKENS must be greater than zero.")
|
raise AuditaConfigError("AUDITA_MAX_SECTION_TOKENS must be greater than zero.")
|
||||||
|
if self.min_section_tokens <= 0:
|
||||||
|
raise AuditaConfigError("AUDITA_MIN_SECTION_TOKENS must be greater than zero.")
|
||||||
|
if self.min_section_tokens > self.max_section_tokens:
|
||||||
|
raise AuditaConfigError(
|
||||||
|
"AUDITA_MIN_SECTION_TOKENS must be less than or equal to AUDITA_MAX_SECTION_TOKENS."
|
||||||
|
)
|
||||||
if not 0.0 <= self.glossary_confidence_threshold <= 1.0:
|
if not 0.0 <= self.glossary_confidence_threshold <= 1.0:
|
||||||
raise AuditaConfigError("AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD must be between 0.0 and 1.0.")
|
raise AuditaConfigError("AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD must be between 0.0 and 1.0.")
|
||||||
if not 0.0 <= self.grammar_confidence_threshold <= 1.0:
|
if not 0.0 <= self.grammar_confidence_threshold <= 1.0:
|
||||||
@@ -217,6 +232,7 @@ class AuditaConfig:
|
|||||||
"base_url": self.base_url,
|
"base_url": self.base_url,
|
||||||
"max_retries": self.max_retries,
|
"max_retries": self.max_retries,
|
||||||
"max_section_tokens": self.max_section_tokens,
|
"max_section_tokens": self.max_section_tokens,
|
||||||
|
"min_section_tokens": self.min_section_tokens,
|
||||||
"glossary_confidence_threshold": self.glossary_confidence_threshold,
|
"glossary_confidence_threshold": self.glossary_confidence_threshold,
|
||||||
"grammar_confidence_threshold": self.grammar_confidence_threshold,
|
"grammar_confidence_threshold": self.grammar_confidence_threshold,
|
||||||
"homophones_confidence_threshold": self.homophones_confidence_threshold,
|
"homophones_confidence_threshold": self.homophones_confidence_threshold,
|
||||||
|
|||||||
17
src/audita/framework/llm_scheduler.py
Normal file
17
src/audita/framework/llm_scheduler.py
Normal file
@@ -0,0 +1,17 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import threading
|
||||||
|
from typing import Callable, TypeVar
|
||||||
|
|
||||||
|
|
||||||
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
|
class ModuleLLMScheduler:
|
||||||
|
def __init__(self, max_concurrency: int) -> None:
|
||||||
|
self.max_concurrency = max_concurrency
|
||||||
|
self._semaphore = threading.BoundedSemaphore(max_concurrency)
|
||||||
|
|
||||||
|
def run_backend_call(self, fn: Callable[[], T]) -> T:
|
||||||
|
with self._semaphore:
|
||||||
|
return fn()
|
||||||
@@ -9,6 +9,8 @@ from audita.core.config import AuditaConfig
|
|||||||
from audita.core.schemas import Glossary
|
from audita.core.schemas import Glossary
|
||||||
from audita.validators.base import Validator
|
from audita.validators.base import Validator
|
||||||
|
|
||||||
|
from .llm_scheduler import ModuleLLMScheduler
|
||||||
|
|
||||||
|
|
||||||
ReplacementPolicy = str
|
ReplacementPolicy = str
|
||||||
|
|
||||||
@@ -47,6 +49,7 @@ class ModuleContext:
|
|||||||
config: AuditaConfig
|
config: AuditaConfig
|
||||||
run_dir: Path
|
run_dir: Path
|
||||||
llm_client: Optional["StructuredLLMClient"] = None
|
llm_client: Optional["StructuredLLMClient"] = None
|
||||||
|
llm_scheduler: Optional[ModuleLLMScheduler] = None
|
||||||
|
|
||||||
|
|
||||||
class StructuredLLMClient(Protocol):
|
class StructuredLLMClient(Protocol):
|
||||||
|
|||||||
@@ -63,12 +63,22 @@ def generate_llm_correction_proposals(
|
|||||||
prompt_path = context.run_dir / f"prompt-{section.section_index:04d}.json"
|
prompt_path = context.run_dir / f"prompt-{section.section_index:04d}.json"
|
||||||
prompt_path.write_text(json.dumps(messages, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
prompt_path.write_text(json.dumps(messages, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||||
|
|
||||||
response = llm_client.run_structured(
|
if context.llm_scheduler is not None:
|
||||||
stage_name=f"{context.run_spec.instance_name}:proposal",
|
response = context.llm_scheduler.run_backend_call(
|
||||||
messages=messages,
|
lambda: llm_client.run_structured(
|
||||||
response_model=StructuredCorrectionSet,
|
stage_name=f"{context.run_spec.instance_name}:proposal",
|
||||||
config=context.config,
|
messages=messages,
|
||||||
)
|
response_model=StructuredCorrectionSet,
|
||||||
|
config=context.config,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
response = llm_client.run_structured(
|
||||||
|
stage_name=f"{context.run_spec.instance_name}:proposal",
|
||||||
|
messages=messages,
|
||||||
|
response_model=StructuredCorrectionSet,
|
||||||
|
config=context.config,
|
||||||
|
)
|
||||||
response_path = context.run_dir / f"corrections-{section.section_index:04d}.json"
|
response_path = context.run_dir / f"corrections-{section.section_index:04d}.json"
|
||||||
response_path.write_text(response.model_dump_json(indent=2) + "\n", encoding="utf-8")
|
response_path.write_text(response.model_dump_json(indent=2) + "\n", encoding="utf-8")
|
||||||
|
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from audita.core.reporting import AppliedChange, ModuleRunReport, ReportedSkip,
|
|||||||
from audita.core.schemas import Glossary, TranscriptSegment
|
from audita.core.schemas import Glossary, TranscriptSegment
|
||||||
from audita.validators.base import ValidationContext, ValidationDecision, ValidationResult
|
from audita.validators.base import ValidationContext, ValidationDecision, ValidationResult
|
||||||
|
|
||||||
|
from .llm_scheduler import ModuleLLMScheduler
|
||||||
from .models import CorrectionProposal, ModuleContext, ModuleRunSpec, StructuredLLMClient
|
from .models import CorrectionProposal, ModuleContext, ModuleRunSpec, StructuredLLMClient
|
||||||
from .proposals import ProposalPreviewError, preview_proposal
|
from .proposals import ProposalPreviewError, preview_proposal
|
||||||
|
|
||||||
@@ -85,8 +86,14 @@ class PipelineRunner:
|
|||||||
config=config,
|
config=config,
|
||||||
run_dir=module_dir,
|
run_dir=module_dir,
|
||||||
llm_client=llm_client,
|
llm_client=llm_client,
|
||||||
|
llm_scheduler=ModuleLLMScheduler(config.llm_concurrency),
|
||||||
|
)
|
||||||
|
sections = chunk_transcript(
|
||||||
|
working,
|
||||||
|
config.max_section_tokens,
|
||||||
|
min_section_tokens=config.min_section_tokens,
|
||||||
|
target_section_count=config.llm_concurrency,
|
||||||
)
|
)
|
||||||
sections = chunk_transcript(working, config.max_section_tokens)
|
|
||||||
if progress is not None:
|
if progress is not None:
|
||||||
progress(
|
progress(
|
||||||
f"Running module {run_spec.instance_name} "
|
f"Running module {run_spec.instance_name} "
|
||||||
@@ -143,6 +150,7 @@ def _run_module(
|
|||||||
llm_client: Optional[StructuredLLMClient],
|
llm_client: Optional[StructuredLLMClient],
|
||||||
) -> _ModuleExecutionResult:
|
) -> _ModuleExecutionResult:
|
||||||
module = context.run_spec.module
|
module = context.run_spec.module
|
||||||
|
validators = list(module.validators())
|
||||||
raw_proposals: List[CorrectionProposal] = []
|
raw_proposals: List[CorrectionProposal] = []
|
||||||
validator_reports: List[ValidatorReport] = []
|
validator_reports: List[ValidatorReport] = []
|
||||||
skipped: List[ReportedSkip] = []
|
skipped: List[ReportedSkip] = []
|
||||||
@@ -167,7 +175,9 @@ def _run_module(
|
|||||||
]
|
]
|
||||||
|
|
||||||
surviving = proposals
|
surviving = proposals
|
||||||
for validator in module.validators():
|
validator_index = 0
|
||||||
|
while validator_index < len(validators):
|
||||||
|
validator = validators[validator_index]
|
||||||
candidate_count = len(surviving)
|
candidate_count = len(surviving)
|
||||||
if not surviving:
|
if not surviving:
|
||||||
validator_reports.append(
|
validator_reports.append(
|
||||||
@@ -179,51 +189,38 @@ def _run_module(
|
|||||||
rejected_count=0,
|
rejected_count=0,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
validator_index += 1
|
||||||
continue
|
continue
|
||||||
|
|
||||||
validation_context = ValidationContext(
|
if validator.execution_kind != "llm":
|
||||||
|
result = _run_single_validator(
|
||||||
|
validator=validator,
|
||||||
|
proposals=surviving,
|
||||||
|
working=working,
|
||||||
|
context=context,
|
||||||
|
llm_client=llm_client,
|
||||||
|
)
|
||||||
|
validator_reports.append(result.report)
|
||||||
|
skipped.extend(result.skipped)
|
||||||
|
surviving = result.approved
|
||||||
|
validator_index += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
group_start = validator_index
|
||||||
|
llm_group = []
|
||||||
|
while validator_index < len(validators) and validators[validator_index].execution_kind == "llm":
|
||||||
|
llm_group.append(validators[validator_index])
|
||||||
|
validator_index += 1
|
||||||
|
group_result = _run_parallel_llm_validator_group(
|
||||||
|
validators=llm_group,
|
||||||
proposals=surviving,
|
proposals=surviving,
|
||||||
transcript=working,
|
working=working,
|
||||||
glossary=context.glossary,
|
context=context,
|
||||||
config=context.config,
|
|
||||||
run_spec=context.run_spec,
|
|
||||||
run_dir=context.run_dir,
|
|
||||||
llm_client=llm_client,
|
llm_client=llm_client,
|
||||||
)
|
)
|
||||||
result = validator.validate(validation_context)
|
validator_reports.extend(group_result.reports)
|
||||||
decisions_by_index = _index_validation_decisions(result, surviving, validator.name)
|
skipped.extend(group_result.skipped)
|
||||||
approved: List[CorrectionProposal] = []
|
surviving = group_result.approved
|
||||||
rejected_count = 0
|
|
||||||
for proposal in surviving:
|
|
||||||
decision = decisions_by_index[proposal.proposal_index]
|
|
||||||
if decision.approved:
|
|
||||||
approved.append(proposal)
|
|
||||||
continue
|
|
||||||
rejected_count += 1
|
|
||||||
skipped.append(
|
|
||||||
ReportedSkip(
|
|
||||||
module_instance=proposal.module_instance,
|
|
||||||
module_key=proposal.module_key,
|
|
||||||
proposal_index=proposal.proposal_index,
|
|
||||||
id=proposal.id,
|
|
||||||
reason=decision.reason or f"{validator.name} rejected proposal",
|
|
||||||
original_text=proposal.original_text,
|
|
||||||
corrected_text=proposal.corrected_text,
|
|
||||||
confidence=proposal.confidence,
|
|
||||||
actual_text=_segment_text_by_id(working).get(proposal.id),
|
|
||||||
source=f"validator:{validator.name}",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
validator_reports.append(
|
|
||||||
ValidatorReport(
|
|
||||||
name=validator.name,
|
|
||||||
execution_kind=validator.execution_kind,
|
|
||||||
candidate_count=candidate_count,
|
|
||||||
approved_count=len(approved),
|
|
||||||
rejected_count=rejected_count,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
surviving = approved
|
|
||||||
|
|
||||||
for proposal in surviving:
|
for proposal in surviving:
|
||||||
apply_result = _apply_proposal(updated_transcript, proposal, module.replacement_policy)
|
apply_result = _apply_proposal(updated_transcript, proposal, module.replacement_policy)
|
||||||
@@ -265,13 +262,161 @@ def _collect_module_proposals(
|
|||||||
context: ModuleContext,
|
context: ModuleContext,
|
||||||
) -> List[List[CorrectionProposal]]:
|
) -> List[List[CorrectionProposal]]:
|
||||||
module = context.run_spec.module
|
module = context.run_spec.module
|
||||||
if context.config.llm_concurrency == 1 or len(sections) <= 1:
|
max_workers = context.llm_scheduler.max_concurrency if context.llm_scheduler is not None else context.config.llm_concurrency
|
||||||
|
if max_workers == 1 or len(sections) <= 1:
|
||||||
return [list(module.propose(section, context)) for section in sections]
|
return [list(module.propose(section, context)) for section in sections]
|
||||||
|
|
||||||
with ThreadPoolExecutor(max_workers=context.config.llm_concurrency) as executor:
|
with ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||||
return list(executor.map(lambda section: list(module.propose(section, context)), sections))
|
return list(executor.map(lambda section: list(module.propose(section, context)), sections))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class _SingleValidatorResult:
|
||||||
|
approved: List[CorrectionProposal]
|
||||||
|
report: ValidatorReport
|
||||||
|
skipped: List[ReportedSkip]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class _ParallelValidatorGroupResult:
|
||||||
|
approved: List[CorrectionProposal]
|
||||||
|
reports: List[ValidatorReport]
|
||||||
|
skipped: List[ReportedSkip]
|
||||||
|
|
||||||
|
|
||||||
|
def _build_validation_context(
|
||||||
|
*,
|
||||||
|
proposals: Sequence[CorrectionProposal],
|
||||||
|
working: Sequence[TranscriptSegment],
|
||||||
|
context: ModuleContext,
|
||||||
|
llm_client: Optional[StructuredLLMClient],
|
||||||
|
) -> ValidationContext:
|
||||||
|
return ValidationContext(
|
||||||
|
proposals=proposals,
|
||||||
|
transcript=working,
|
||||||
|
glossary=context.glossary,
|
||||||
|
config=context.config,
|
||||||
|
run_spec=context.run_spec,
|
||||||
|
run_dir=context.run_dir,
|
||||||
|
llm_client=llm_client,
|
||||||
|
llm_scheduler=context.llm_scheduler,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _run_single_validator(
|
||||||
|
*,
|
||||||
|
validator,
|
||||||
|
proposals: Sequence[CorrectionProposal],
|
||||||
|
working: Sequence[TranscriptSegment],
|
||||||
|
context: ModuleContext,
|
||||||
|
llm_client: Optional[StructuredLLMClient],
|
||||||
|
) -> _SingleValidatorResult:
|
||||||
|
candidate_count = len(proposals)
|
||||||
|
validation_context = _build_validation_context(
|
||||||
|
proposals=proposals,
|
||||||
|
working=working,
|
||||||
|
context=context,
|
||||||
|
llm_client=llm_client,
|
||||||
|
)
|
||||||
|
result = validator.validate(validation_context)
|
||||||
|
decisions_by_index = _index_validation_decisions(result, proposals, validator.name)
|
||||||
|
approved: List[CorrectionProposal] = []
|
||||||
|
skipped: List[ReportedSkip] = []
|
||||||
|
rejected_count = 0
|
||||||
|
for proposal in proposals:
|
||||||
|
decision = decisions_by_index[proposal.proposal_index]
|
||||||
|
if decision.approved:
|
||||||
|
approved.append(proposal)
|
||||||
|
continue
|
||||||
|
rejected_count += 1
|
||||||
|
skipped.append(_reported_skip_from_decision(proposal, working, validator.name, decision))
|
||||||
|
return _SingleValidatorResult(
|
||||||
|
approved=approved,
|
||||||
|
report=ValidatorReport(
|
||||||
|
name=validator.name,
|
||||||
|
execution_kind=validator.execution_kind,
|
||||||
|
candidate_count=candidate_count,
|
||||||
|
approved_count=len(approved),
|
||||||
|
rejected_count=rejected_count,
|
||||||
|
),
|
||||||
|
skipped=skipped,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _run_parallel_llm_validator_group(
|
||||||
|
*,
|
||||||
|
validators: Sequence,
|
||||||
|
proposals: Sequence[CorrectionProposal],
|
||||||
|
working: Sequence[TranscriptSegment],
|
||||||
|
context: ModuleContext,
|
||||||
|
llm_client: Optional[StructuredLLMClient],
|
||||||
|
) -> _ParallelValidatorGroupResult:
|
||||||
|
validation_context = _build_validation_context(
|
||||||
|
proposals=proposals,
|
||||||
|
working=working,
|
||||||
|
context=context,
|
||||||
|
llm_client=llm_client,
|
||||||
|
)
|
||||||
|
max_workers = context.llm_scheduler.max_concurrency if context.llm_scheduler is not None else context.config.llm_concurrency
|
||||||
|
if max_workers == 1 or len(validators) <= 1:
|
||||||
|
results = [validator.validate(validation_context) for validator in validators]
|
||||||
|
else:
|
||||||
|
with ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||||
|
results = list(executor.map(lambda validator: validator.validate(validation_context), validators))
|
||||||
|
|
||||||
|
indexed_results = [
|
||||||
|
_index_validation_decisions(result, proposals, validator.name)
|
||||||
|
for validator, result in zip(validators, results)
|
||||||
|
]
|
||||||
|
reports = [
|
||||||
|
ValidatorReport(
|
||||||
|
name=validator.name,
|
||||||
|
execution_kind=validator.execution_kind,
|
||||||
|
candidate_count=len(proposals),
|
||||||
|
approved_count=sum(1 for decision in decisions.values() if decision.approved),
|
||||||
|
rejected_count=sum(1 for decision in decisions.values() if not decision.approved),
|
||||||
|
)
|
||||||
|
for validator, decisions in zip(validators, indexed_results)
|
||||||
|
]
|
||||||
|
|
||||||
|
approved: List[CorrectionProposal] = []
|
||||||
|
skipped: List[ReportedSkip] = []
|
||||||
|
for proposal in proposals:
|
||||||
|
rejection = None
|
||||||
|
for validator, decisions in zip(validators, indexed_results):
|
||||||
|
decision = decisions[proposal.proposal_index]
|
||||||
|
if not decision.approved:
|
||||||
|
rejection = (validator.name, decision)
|
||||||
|
break
|
||||||
|
if rejection is None:
|
||||||
|
approved.append(proposal)
|
||||||
|
continue
|
||||||
|
validator_name, decision = rejection
|
||||||
|
skipped.append(_reported_skip_from_decision(proposal, working, validator_name, decision))
|
||||||
|
|
||||||
|
return _ParallelValidatorGroupResult(approved=approved, reports=reports, skipped=skipped)
|
||||||
|
|
||||||
|
|
||||||
|
def _reported_skip_from_decision(
|
||||||
|
proposal: CorrectionProposal,
|
||||||
|
working: Sequence[TranscriptSegment],
|
||||||
|
validator_name: str,
|
||||||
|
decision: ValidationDecision,
|
||||||
|
) -> ReportedSkip:
|
||||||
|
return ReportedSkip(
|
||||||
|
module_instance=proposal.module_instance,
|
||||||
|
module_key=proposal.module_key,
|
||||||
|
proposal_index=proposal.proposal_index,
|
||||||
|
id=proposal.id,
|
||||||
|
reason=decision.reason or f"{validator_name} rejected proposal",
|
||||||
|
original_text=proposal.original_text,
|
||||||
|
corrected_text=proposal.corrected_text,
|
||||||
|
confidence=proposal.confidence,
|
||||||
|
actual_text=_segment_text_by_id(working).get(proposal.id),
|
||||||
|
source=f"validator:{validator_name}",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _index_validation_decisions(
|
def _index_validation_decisions(
|
||||||
result: ValidationResult,
|
result: ValidationResult,
|
||||||
proposals: Sequence[CorrectionProposal],
|
proposals: Sequence[CorrectionProposal],
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from audita.core.schemas import Glossary, TranscriptSegment
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from audita.framework.models import CorrectionProposal, ModuleRunSpec, StructuredLLMClient
|
from audita.framework.models import CorrectionProposal, ModuleRunSpec, StructuredLLMClient
|
||||||
|
from audita.framework.llm_scheduler import ModuleLLMScheduler
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -20,6 +21,7 @@ class ValidationContext:
|
|||||||
run_spec: "ModuleRunSpec"
|
run_spec: "ModuleRunSpec"
|
||||||
run_dir: Path
|
run_dir: Path
|
||||||
llm_client: Optional["StructuredLLMClient"] = None
|
llm_client: Optional["StructuredLLMClient"] = None
|
||||||
|
llm_scheduler: Optional["ModuleLLMScheduler"] = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
|
|||||||
@@ -79,11 +79,12 @@ class _BaseLLMValidator:
|
|||||||
payload_fn=lambda item: item.to_prompt_payload(),
|
payload_fn=lambda item: item.to_prompt_payload(),
|
||||||
empty_error_message="Validation input must contain at least one proposal.",
|
empty_error_message="Validation input must contain at least one proposal.",
|
||||||
)
|
)
|
||||||
if context.config.llm_concurrency == 1 or len(batches) <= 1:
|
max_workers = context.llm_scheduler.max_concurrency if context.llm_scheduler is not None else context.config.llm_concurrency
|
||||||
|
if max_workers == 1 or len(batches) <= 1:
|
||||||
for batch in batches:
|
for batch in batches:
|
||||||
decisions.extend(self._run_batch(context, llm_client, batch))
|
decisions.extend(self._run_batch(context, llm_client, batch))
|
||||||
else:
|
else:
|
||||||
with ThreadPoolExecutor(max_workers=context.config.llm_concurrency) as executor:
|
with ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||||
for batch_decisions in executor.map(
|
for batch_decisions in executor.map(
|
||||||
lambda batch: self._run_batch(context, llm_client, batch),
|
lambda batch: self._run_batch(context, llm_client, batch),
|
||||||
batches,
|
batches,
|
||||||
@@ -141,12 +142,22 @@ class _BaseLLMValidator:
|
|||||||
prompt_path = context.run_dir / f"{self.name}-prompt-{batch.batch_index:04d}.json"
|
prompt_path = context.run_dir / f"{self.name}-prompt-{batch.batch_index:04d}.json"
|
||||||
response_path = context.run_dir / f"{self.name}-response-{batch.batch_index:04d}.json"
|
response_path = context.run_dir / f"{self.name}-response-{batch.batch_index:04d}.json"
|
||||||
_write_json(prompt_path, {"messages": messages})
|
_write_json(prompt_path, {"messages": messages})
|
||||||
response = llm_client.run_structured(
|
if context.llm_scheduler is not None:
|
||||||
stage_name=f"{context.run_spec.instance_name}:{self.name}",
|
response = context.llm_scheduler.run_backend_call(
|
||||||
messages=messages,
|
lambda: llm_client.run_structured(
|
||||||
response_model=_LLMValidationSetModel,
|
stage_name=f"{context.run_spec.instance_name}:{self.name}",
|
||||||
config=context.config,
|
messages=messages,
|
||||||
)
|
response_model=_LLMValidationSetModel,
|
||||||
|
config=context.config,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
response = llm_client.run_structured(
|
||||||
|
stage_name=f"{context.run_spec.instance_name}:{self.name}",
|
||||||
|
messages=messages,
|
||||||
|
response_model=_LLMValidationSetModel,
|
||||||
|
config=context.config,
|
||||||
|
)
|
||||||
_write_json(response_path, response.model_dump(mode="json"))
|
_write_json(response_path, response.model_dump(mode="json"))
|
||||||
return self._validate_batch_response(response, batch.items)
|
return self._validate_batch_response(response, batch.items)
|
||||||
|
|
||||||
|
|||||||
@@ -27,6 +27,79 @@ def test_chunk_transcript_batches_sections_by_token_limit():
|
|||||||
assert [segment.segment.id for segment in sections[1].segments] == [3]
|
assert [segment.segment.id for segment in sections[1].segments] == [3]
|
||||||
|
|
||||||
|
|
||||||
|
def test_chunk_transcript_targets_llm_concurrency_when_feasible():
|
||||||
|
transcript = parse_transcript_json(
|
||||||
|
"""
|
||||||
|
[
|
||||||
|
{"id": 1, "speaker": "A", "start": 0.0, "end": 1.0, "text": "one"},
|
||||||
|
{"id": 2, "speaker": "A", "start": 1.0, "end": 2.0, "text": "two"},
|
||||||
|
{"id": 3, "speaker": "A", "start": 2.0, "end": 3.0, "text": "three"},
|
||||||
|
{"id": 4, "speaker": "A", "start": 3.0, "end": 4.0, "text": "four"}
|
||||||
|
]
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
sections = chunk_transcript(
|
||||||
|
transcript,
|
||||||
|
max_section_tokens=8,
|
||||||
|
min_section_tokens=4,
|
||||||
|
target_section_count=2,
|
||||||
|
estimator=FakeEstimator(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(sections) == 2
|
||||||
|
assert [segment.segment.id for segment in sections[0].segments] == [1, 2]
|
||||||
|
assert [segment.segment.id for segment in sections[1].segments] == [3, 4]
|
||||||
|
|
||||||
|
|
||||||
|
def test_chunk_transcript_increases_section_count_when_target_sections_exceed_max_tokens():
|
||||||
|
transcript = parse_transcript_json(
|
||||||
|
"""
|
||||||
|
[
|
||||||
|
{"id": 1, "speaker": "A", "start": 0.0, "end": 1.0, "text": "one"},
|
||||||
|
{"id": 2, "speaker": "A", "start": 1.0, "end": 2.0, "text": "two"},
|
||||||
|
{"id": 3, "speaker": "A", "start": 2.0, "end": 3.0, "text": "three"},
|
||||||
|
{"id": 4, "speaker": "A", "start": 3.0, "end": 4.0, "text": "four"},
|
||||||
|
{"id": 5, "speaker": "A", "start": 4.0, "end": 5.0, "text": "five"}
|
||||||
|
]
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
sections = chunk_transcript(
|
||||||
|
transcript,
|
||||||
|
max_section_tokens=8,
|
||||||
|
min_section_tokens=4,
|
||||||
|
target_section_count=1,
|
||||||
|
estimator=FakeEstimator(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(sections) > 1
|
||||||
|
assert all(section.token_count <= 8 for section in sections)
|
||||||
|
|
||||||
|
|
||||||
|
def test_chunk_transcript_reduces_section_count_when_target_sections_fall_below_min_tokens():
|
||||||
|
transcript = parse_transcript_json(
|
||||||
|
"""
|
||||||
|
[
|
||||||
|
{"id": 1, "speaker": "A", "start": 0.0, "end": 1.0, "text": "one"},
|
||||||
|
{"id": 2, "speaker": "A", "start": 1.0, "end": 2.0, "text": "two"},
|
||||||
|
{"id": 3, "speaker": "A", "start": 2.0, "end": 3.0, "text": "three"}
|
||||||
|
]
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
sections = chunk_transcript(
|
||||||
|
transcript,
|
||||||
|
max_section_tokens=12,
|
||||||
|
min_section_tokens=8,
|
||||||
|
target_section_count=3,
|
||||||
|
estimator=FakeEstimator(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(sections) == 1
|
||||||
|
assert sections[0].token_count >= 8
|
||||||
|
|
||||||
|
|
||||||
def test_chunk_transcript_prompt_payload_includes_categories_when_present():
|
def test_chunk_transcript_prompt_payload_includes_categories_when_present():
|
||||||
transcript = parse_transcript_json(
|
transcript = parse_transcript_json(
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -45,23 +45,70 @@ class RecordingLLMValidator(RecordingValidator):
|
|||||||
|
|
||||||
class FakeStructuredLLMClient:
|
class FakeStructuredLLMClient:
|
||||||
def __init__(self, responses):
|
def __init__(self, responses):
|
||||||
self._responses = list(responses)
|
self._responses = responses
|
||||||
|
self._lock = threading.Lock()
|
||||||
self.calls = []
|
self.calls = []
|
||||||
|
|
||||||
def run_structured(self, *, stage_name, messages, response_model, config):
|
def run_structured(self, *, stage_name, messages, response_model, config):
|
||||||
self.calls.append(
|
with self._lock:
|
||||||
{
|
self.calls.append(
|
||||||
"stage_name": stage_name,
|
{
|
||||||
"messages": list(messages),
|
"stage_name": stage_name,
|
||||||
"response_model": response_model,
|
"messages": list(messages),
|
||||||
}
|
"response_model": response_model,
|
||||||
)
|
}
|
||||||
if not self._responses:
|
)
|
||||||
raise AuditaLLMError("FakeStructuredLLMClient received more calls than expected.")
|
payload = _pop_llm_response(self._responses, stage_name)
|
||||||
payload = self._responses.pop(0)
|
|
||||||
return response_model.model_validate(payload)
|
return response_model.model_validate(payload)
|
||||||
|
|
||||||
|
|
||||||
|
def _pop_llm_response(responses, stage_name):
|
||||||
|
if isinstance(responses, dict):
|
||||||
|
if stage_name not in responses:
|
||||||
|
raise AuditaLLMError(f"FakeStructuredLLMClient received unexpected stage_name: {stage_name}")
|
||||||
|
payloads = responses[stage_name]
|
||||||
|
if isinstance(payloads, list):
|
||||||
|
if not payloads:
|
||||||
|
raise AuditaLLMError(f"FakeStructuredLLMClient received too many calls for stage_name: {stage_name}")
|
||||||
|
return payloads.pop(0)
|
||||||
|
payload = payloads
|
||||||
|
del responses[stage_name]
|
||||||
|
return payload
|
||||||
|
if not responses:
|
||||||
|
raise AuditaLLMError("FakeStructuredLLMClient received more calls than expected.")
|
||||||
|
return responses.pop(0)
|
||||||
|
|
||||||
|
|
||||||
|
class TrackingStructuredLLMClient:
|
||||||
|
def __init__(self, responses, barrier=None):
|
||||||
|
self._responses = responses
|
||||||
|
self._barrier = barrier
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
self.calls = []
|
||||||
|
self.in_flight = 0
|
||||||
|
self.max_in_flight = 0
|
||||||
|
|
||||||
|
def run_structured(self, *, stage_name, messages, response_model, config):
|
||||||
|
if self._barrier is not None:
|
||||||
|
self._barrier.wait()
|
||||||
|
with self._lock:
|
||||||
|
self.in_flight += 1
|
||||||
|
self.max_in_flight = max(self.max_in_flight, self.in_flight)
|
||||||
|
self.calls.append(
|
||||||
|
{
|
||||||
|
"stage_name": stage_name,
|
||||||
|
"messages": list(messages),
|
||||||
|
"response_model": response_model,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
payload = _pop_llm_response(self._responses, stage_name)
|
||||||
|
try:
|
||||||
|
return response_model.model_validate(payload)
|
||||||
|
finally:
|
||||||
|
with self._lock:
|
||||||
|
self.in_flight -= 1
|
||||||
|
|
||||||
|
|
||||||
class RecordingModule:
|
class RecordingModule:
|
||||||
replacement_policy = "require_unique"
|
replacement_policy = "require_unique"
|
||||||
|
|
||||||
@@ -322,8 +369,8 @@ def test_pipeline_runner_supports_real_llm_validators_in_one_chain(tmp_path):
|
|||||||
[],
|
[],
|
||||||
)
|
)
|
||||||
llm_client = FakeStructuredLLMClient(
|
llm_client = FakeStructuredLLMClient(
|
||||||
[
|
{
|
||||||
{
|
"mixed_real:spoken_form_plausibility_review": {
|
||||||
"validations": [
|
"validations": [
|
||||||
{
|
{
|
||||||
"correction_index": 0,
|
"correction_index": 0,
|
||||||
@@ -333,7 +380,7 @@ def test_pipeline_runner_supports_real_llm_validators_in_one_chain(tmp_path):
|
|||||||
}
|
}
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
"mixed_real:meaning_reversal_review": {
|
||||||
"validations": [
|
"validations": [
|
||||||
{
|
{
|
||||||
"correction_index": 0,
|
"correction_index": 0,
|
||||||
@@ -343,7 +390,7 @@ def test_pipeline_runner_supports_real_llm_validators_in_one_chain(tmp_path):
|
|||||||
}
|
}
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
]
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
runner = PipelineRunner()
|
runner = PipelineRunner()
|
||||||
@@ -363,10 +410,10 @@ def test_pipeline_runner_supports_real_llm_validators_in_one_chain(tmp_path):
|
|||||||
"llm",
|
"llm",
|
||||||
"llm",
|
"llm",
|
||||||
]
|
]
|
||||||
assert [call["stage_name"] for call in llm_client.calls] == [
|
assert {call["stage_name"] for call in llm_client.calls} == {
|
||||||
"mixed_real:spoken_form_plausibility_review",
|
"mixed_real:spoken_form_plausibility_review",
|
||||||
"mixed_real:meaning_reversal_review",
|
"mixed_real:meaning_reversal_review",
|
||||||
]
|
}
|
||||||
|
|
||||||
|
|
||||||
def test_pipeline_runner_uses_real_protected_glossary_validator(tmp_path):
|
def test_pipeline_runner_uses_real_protected_glossary_validator(tmp_path):
|
||||||
@@ -449,7 +496,10 @@ def test_pipeline_runner_collects_section_proposals_concurrently_and_preserves_s
|
|||||||
]
|
]
|
||||||
seen = []
|
seen = []
|
||||||
module = ConcurrentRecordingModule(seen, threading.Barrier(2, timeout=1.0))
|
module = ConcurrentRecordingModule(seen, threading.Barrier(2, timeout=1.0))
|
||||||
monkeypatch.setattr("audita.framework.runner.chunk_transcript", lambda working, max_tokens: sections)
|
monkeypatch.setattr(
|
||||||
|
"audita.framework.runner.chunk_transcript",
|
||||||
|
lambda working, max_tokens, min_section_tokens=1, target_section_count=None: sections,
|
||||||
|
)
|
||||||
|
|
||||||
runner = PipelineRunner()
|
runner = PipelineRunner()
|
||||||
result = runner.run(
|
result = runner.run(
|
||||||
@@ -465,3 +515,79 @@ def test_pipeline_runner_collects_section_proposals_concurrently_and_preserves_s
|
|||||||
assert result.applied_changes[1].corrected_text == "Beta revised"
|
assert result.applied_changes[1].corrected_text == "Beta revised"
|
||||||
assert result.transcript[0].text == "Alpha revised."
|
assert result.transcript[0].text == "Alpha revised."
|
||||||
assert result.transcript[1].text == "Beta revised."
|
assert result.transcript[1].text == "Beta revised."
|
||||||
|
|
||||||
|
|
||||||
|
def test_pipeline_runner_reports_first_llm_validator_rejection_in_chain_order(tmp_path):
|
||||||
|
transcript = parse_transcript_json(
|
||||||
|
"""
|
||||||
|
[
|
||||||
|
{"id": 1, "speaker": "A", "start": 0.0, "end": 1.0, "text": "There were gestures at the dam."}
|
||||||
|
]
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
glossary = parse_glossary_yaml(
|
||||||
|
"""
|
||||||
|
glossary:
|
||||||
|
- name: "Jesters"
|
||||||
|
category: faction
|
||||||
|
summary: "The Jesters are a faction."
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
module = RecordingModule(
|
||||||
|
"ordered_real",
|
||||||
|
[
|
||||||
|
CorrectionProposal(
|
||||||
|
proposal_index=0,
|
||||||
|
module_instance="ordered_real",
|
||||||
|
module_key="glossary",
|
||||||
|
id=1,
|
||||||
|
original_text="gestures",
|
||||||
|
corrected_text="Jesters",
|
||||||
|
confidence=0.9,
|
||||||
|
)
|
||||||
|
],
|
||||||
|
[
|
||||||
|
ProposalConfidenceValidator("proposal_confidence_guard", "glossary_confidence_threshold"),
|
||||||
|
ProtectedGlossaryTermsValidator("protected_glossary_guard"),
|
||||||
|
SpokenFormPlausibilityValidator("spoken_form_plausibility_review"),
|
||||||
|
MeaningReversalValidator("meaning_reversal_review"),
|
||||||
|
],
|
||||||
|
[],
|
||||||
|
)
|
||||||
|
llm_client = FakeStructuredLLMClient(
|
||||||
|
{
|
||||||
|
"ordered_real:spoken_form_plausibility_review": {
|
||||||
|
"validations": [
|
||||||
|
{
|
||||||
|
"correction_index": 0,
|
||||||
|
"approved": False,
|
||||||
|
"confidence": 0.95,
|
||||||
|
"reason": "Not plausibly supported by spoken-form context.",
|
||||||
|
}
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"ordered_real:meaning_reversal_review": {
|
||||||
|
"validations": [
|
||||||
|
{
|
||||||
|
"correction_index": 0,
|
||||||
|
"approved": False,
|
||||||
|
"confidence": 0.98,
|
||||||
|
"reason": "Changes meaning too much.",
|
||||||
|
}
|
||||||
|
]
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
runner = PipelineRunner()
|
||||||
|
result = runner.run(
|
||||||
|
transcript=transcript,
|
||||||
|
glossary=glossary,
|
||||||
|
module_specs=[ModuleRunSpec(instance_name="ordered_real", module_key="glossary", module=module)],
|
||||||
|
config=AuditaConfig.from_sources(env={"OPENROUTER_API_KEY": "test-key"}),
|
||||||
|
run_dir=tmp_path / "run",
|
||||||
|
llm_client=llm_client,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.skipped_corrections[0].source == "validator:spoken_form_plausibility_review"
|
||||||
|
assert result.skipped_corrections[0].reason == "Not plausibly supported by spoken-form context."
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import threading
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from audita.core.chunking import chunk_transcript
|
from audita.core.chunking import chunk_transcript
|
||||||
@@ -19,20 +20,38 @@ from audita.pipeline import process_transcript_result
|
|||||||
|
|
||||||
class FakeStructuredLLMClient:
|
class FakeStructuredLLMClient:
|
||||||
def __init__(self, responses):
|
def __init__(self, responses):
|
||||||
self._responses = list(responses)
|
self._responses = responses
|
||||||
|
self._lock = threading.Lock()
|
||||||
self.calls = []
|
self.calls = []
|
||||||
|
|
||||||
def run_structured(self, *, stage_name, messages, response_model, config):
|
def run_structured(self, *, stage_name, messages, response_model, config):
|
||||||
self.calls.append(
|
with self._lock:
|
||||||
{
|
self.calls.append(
|
||||||
"stage_name": stage_name,
|
{
|
||||||
"messages": list(messages),
|
"stage_name": stage_name,
|
||||||
"response_model": response_model,
|
"messages": list(messages),
|
||||||
}
|
"response_model": response_model,
|
||||||
)
|
}
|
||||||
if not self._responses:
|
)
|
||||||
raise AuditaLLMError("FakeStructuredLLMClient received more calls than expected.")
|
payload = _pop_llm_response(self._responses, stage_name)
|
||||||
return response_model.model_validate(self._responses.pop(0))
|
return response_model.model_validate(payload)
|
||||||
|
|
||||||
|
|
||||||
|
def _pop_llm_response(responses, stage_name):
|
||||||
|
if isinstance(responses, dict):
|
||||||
|
if stage_name not in responses:
|
||||||
|
raise AuditaLLMError(f"FakeStructuredLLMClient received unexpected stage_name: {stage_name}")
|
||||||
|
payloads = responses[stage_name]
|
||||||
|
if isinstance(payloads, list):
|
||||||
|
if not payloads:
|
||||||
|
raise AuditaLLMError(f"FakeStructuredLLMClient received too many calls for stage_name: {stage_name}")
|
||||||
|
return payloads.pop(0)
|
||||||
|
payload = payloads
|
||||||
|
del responses[stage_name]
|
||||||
|
return payload
|
||||||
|
if not responses:
|
||||||
|
raise AuditaLLMError("FakeStructuredLLMClient received more calls than expected.")
|
||||||
|
return responses.pop(0)
|
||||||
|
|
||||||
|
|
||||||
def _glossary():
|
def _glossary():
|
||||||
@@ -810,8 +829,8 @@ def test_process_transcript_result_grammar_module_still_rejects_homophone_style_
|
|||||||
work_dir_retention="always",
|
work_dir_retention="always",
|
||||||
)
|
)
|
||||||
client = FakeStructuredLLMClient(
|
client = FakeStructuredLLMClient(
|
||||||
[
|
{
|
||||||
{
|
"grammar:proposal": {
|
||||||
"corrections": [
|
"corrections": [
|
||||||
{
|
{
|
||||||
"id": 1,
|
"id": 1,
|
||||||
@@ -821,7 +840,7 @@ def test_process_transcript_result_grammar_module_still_rejects_homophone_style_
|
|||||||
}
|
}
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
"grammar:grammar_only_guard": {
|
||||||
"validations": [
|
"validations": [
|
||||||
{
|
{
|
||||||
"correction_index": 0,
|
"correction_index": 0,
|
||||||
@@ -830,8 +849,18 @@ def test_process_transcript_result_grammar_module_still_rejects_homophone_style_
|
|||||||
"reason": "Free-standing homophone rewrite rather than conservative grammar cleanup.",
|
"reason": "Free-standing homophone rewrite rather than conservative grammar cleanup.",
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
},
|
||||||
]
|
"grammar:meaning_reversal_review": {
|
||||||
|
"validations": [
|
||||||
|
{
|
||||||
|
"correction_index": 0,
|
||||||
|
"approved": True,
|
||||||
|
"confidence": 0.99,
|
||||||
|
"reason": "Does not reverse the segment meaning.",
|
||||||
|
}
|
||||||
|
]
|
||||||
|
},
|
||||||
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
result = process_transcript_result(
|
result = process_transcript_result(
|
||||||
@@ -846,6 +875,7 @@ def test_process_transcript_result_grammar_module_still_rejects_homophone_style_
|
|||||||
assert [call["stage_name"] for call in client.calls] == [
|
assert [call["stage_name"] for call in client.calls] == [
|
||||||
"grammar:proposal",
|
"grammar:proposal",
|
||||||
"grammar:grammar_only_guard",
|
"grammar:grammar_only_guard",
|
||||||
|
"grammar:meaning_reversal_review",
|
||||||
]
|
]
|
||||||
assert result.report.skipped_corrections[0].source == "validator:grammar_only_guard"
|
assert result.report.skipped_corrections[0].source == "validator:grammar_only_guard"
|
||||||
assert "grammar cleanup" in result.report.skipped_corrections[0].reason
|
assert "grammar cleanup" in result.report.skipped_corrections[0].reason
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ def test_process_help_exposes_framework_flags(capsys):
|
|||||||
assert "--base-url" in output
|
assert "--base-url" in output
|
||||||
assert "--max-retries" in output
|
assert "--max-retries" in output
|
||||||
assert "--max-section-tokens" in output
|
assert "--max-section-tokens" in output
|
||||||
|
assert "--min-section-tokens" in output
|
||||||
assert "--glossary-confidence-threshold" in output
|
assert "--glossary-confidence-threshold" in output
|
||||||
assert "--grammar-confidence-threshold" in output
|
assert "--grammar-confidence-threshold" in output
|
||||||
assert "--homophones-confidence-threshold" in output
|
assert "--homophones-confidence-threshold" in output
|
||||||
@@ -249,3 +250,57 @@ def test_cli_process_passes_llm_concurrency_override_to_config(monkeypatch, tmp_
|
|||||||
|
|
||||||
assert exit_code == 0
|
assert exit_code == 0
|
||||||
assert captured["llm_concurrency"] == 3
|
assert captured["llm_concurrency"] == 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_cli_process_passes_min_section_tokens_override_to_config(monkeypatch, tmp_path):
|
||||||
|
captured = {}
|
||||||
|
transcript = parse_transcript_json(
|
||||||
|
"""
|
||||||
|
[
|
||||||
|
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Fixed."}
|
||||||
|
]
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
report = RunReport(
|
||||||
|
status="success",
|
||||||
|
config={"model": "m", "base_url": "b"},
|
||||||
|
normalization={"source_segment_count": 1, "normalized_segment_count": 1, "merge_count": 0},
|
||||||
|
pipeline=["grammar"],
|
||||||
|
modules=[],
|
||||||
|
applied_changes=[],
|
||||||
|
skipped_corrections=[],
|
||||||
|
totals={"output_segment_count": 1, "applied_change_count": 0, "skipped_correction_count": 0},
|
||||||
|
work_dir_retention="auto",
|
||||||
|
work_dir_retained=False,
|
||||||
|
work_dir=None,
|
||||||
|
error=None,
|
||||||
|
)
|
||||||
|
result = ProcessResult(
|
||||||
|
transcript=transcript,
|
||||||
|
report=report,
|
||||||
|
run_dir=tmp_path / "run",
|
||||||
|
work_dir_retained=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _fake_from_sources(*, overrides=None):
|
||||||
|
captured["min_section_tokens"] = overrides.min_section_tokens
|
||||||
|
return object()
|
||||||
|
|
||||||
|
monkeypatch.setattr("audita.cli.AuditaConfig.from_sources", _fake_from_sources)
|
||||||
|
monkeypatch.setattr("audita.cli.load_transcript", lambda path: [])
|
||||||
|
monkeypatch.setattr("audita.cli.load_glossary", lambda path: object())
|
||||||
|
monkeypatch.setattr("audita.cli.process_transcript_result", lambda *args, **kwargs: result)
|
||||||
|
|
||||||
|
exit_code = main(
|
||||||
|
[
|
||||||
|
"process",
|
||||||
|
"transcript.json",
|
||||||
|
"--glossary",
|
||||||
|
"glossary.yaml",
|
||||||
|
"--min-section-tokens",
|
||||||
|
"5000",
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert exit_code == 0
|
||||||
|
assert captured["min_section_tokens"] == 5000
|
||||||
|
|||||||
@@ -7,6 +7,8 @@ from audita.core.config import (
|
|||||||
DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD,
|
DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD,
|
||||||
DEFAULT_HOMOPHONES_CONFIDENCE_THRESHOLD,
|
DEFAULT_HOMOPHONES_CONFIDENCE_THRESHOLD,
|
||||||
DEFAULT_LLM_CONCURRENCY,
|
DEFAULT_LLM_CONCURRENCY,
|
||||||
|
DEFAULT_MAX_SECTION_TOKENS,
|
||||||
|
DEFAULT_MIN_SECTION_TOKENS,
|
||||||
DEFAULT_NORMALIZE_MAX_SEGMENT_GAP,
|
DEFAULT_NORMALIZE_MAX_SEGMENT_GAP,
|
||||||
DEFAULT_SPOKEN_WORD_CONFIDENCE_THRESHOLD,
|
DEFAULT_SPOKEN_WORD_CONFIDENCE_THRESHOLD,
|
||||||
DEFAULT_WORK_DIR_RETENTION,
|
DEFAULT_WORK_DIR_RETENTION,
|
||||||
@@ -20,6 +22,8 @@ def test_default_config_allows_missing_api_key():
|
|||||||
|
|
||||||
assert config.api_key is None
|
assert config.api_key is None
|
||||||
assert config.llm_concurrency == DEFAULT_LLM_CONCURRENCY
|
assert config.llm_concurrency == DEFAULT_LLM_CONCURRENCY
|
||||||
|
assert config.max_section_tokens == DEFAULT_MAX_SECTION_TOKENS
|
||||||
|
assert config.min_section_tokens == DEFAULT_MIN_SECTION_TOKENS
|
||||||
assert config.module_keys == DEFAULT_MODULE_KEYS
|
assert config.module_keys == DEFAULT_MODULE_KEYS
|
||||||
assert config.glossary_confidence_threshold == DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD
|
assert config.glossary_confidence_threshold == DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD
|
||||||
assert config.grammar_confidence_threshold == DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD
|
assert config.grammar_confidence_threshold == DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD
|
||||||
@@ -32,12 +36,21 @@ def test_default_config_allows_missing_api_key():
|
|||||||
def test_cli_overrides_take_precedence():
|
def test_cli_overrides_take_precedence():
|
||||||
config = AuditaConfig.from_sources(
|
config = AuditaConfig.from_sources(
|
||||||
env={"AUDITA_MAX_SECTION_TOKENS": "1000"},
|
env={"AUDITA_MAX_SECTION_TOKENS": "1000"},
|
||||||
overrides=ConfigOverrides(max_section_tokens=2000),
|
overrides=ConfigOverrides(max_section_tokens=2000, min_section_tokens=1000),
|
||||||
)
|
)
|
||||||
|
|
||||||
assert config.max_section_tokens == 2000
|
assert config.max_section_tokens == 2000
|
||||||
|
|
||||||
|
|
||||||
|
def test_min_section_tokens_cli_override_takes_precedence():
|
||||||
|
config = AuditaConfig.from_sources(
|
||||||
|
env={"AUDITA_MIN_SECTION_TOKENS": "2000"},
|
||||||
|
overrides=ConfigOverrides(min_section_tokens=6000),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert config.min_section_tokens == 6000
|
||||||
|
|
||||||
|
|
||||||
def test_llm_concurrency_cli_override_takes_precedence():
|
def test_llm_concurrency_cli_override_takes_precedence():
|
||||||
config = AuditaConfig.from_sources(
|
config = AuditaConfig.from_sources(
|
||||||
env={"AUDITA_LLM_CONCURRENCY": "2"},
|
env={"AUDITA_LLM_CONCURRENCY": "2"},
|
||||||
@@ -53,6 +66,12 @@ def test_llm_concurrency_env_is_parsed():
|
|||||||
assert config.llm_concurrency == 3
|
assert config.llm_concurrency == 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_min_section_tokens_env_is_parsed():
|
||||||
|
config = AuditaConfig.from_sources(env={"AUDITA_MIN_SECTION_TOKENS": "3000"})
|
||||||
|
|
||||||
|
assert config.min_section_tokens == 3000
|
||||||
|
|
||||||
|
|
||||||
def test_generic_llm_api_key_env_is_read():
|
def test_generic_llm_api_key_env_is_read():
|
||||||
config = AuditaConfig.from_sources(env={"AUDITA_LLM_API_KEY": "generic-key"})
|
config = AuditaConfig.from_sources(env={"AUDITA_LLM_API_KEY": "generic-key"})
|
||||||
|
|
||||||
@@ -164,3 +183,19 @@ def test_invalid_thresholds_are_rejected(env_name):
|
|||||||
def test_invalid_llm_concurrency_is_rejected(value):
|
def test_invalid_llm_concurrency_is_rejected(value):
|
||||||
with pytest.raises(AuditaConfigError, match="AUDITA_LLM_CONCURRENCY"):
|
with pytest.raises(AuditaConfigError, match="AUDITA_LLM_CONCURRENCY"):
|
||||||
AuditaConfig.from_sources(env={"AUDITA_LLM_CONCURRENCY": value})
|
AuditaConfig.from_sources(env={"AUDITA_LLM_CONCURRENCY": value})
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("value", ["0", "-1", "many"])
|
||||||
|
def test_invalid_min_section_tokens_is_rejected(value):
|
||||||
|
with pytest.raises(AuditaConfigError, match="AUDITA_MIN_SECTION_TOKENS"):
|
||||||
|
AuditaConfig.from_sources(env={"AUDITA_MIN_SECTION_TOKENS": value})
|
||||||
|
|
||||||
|
|
||||||
|
def test_min_section_tokens_must_not_exceed_max_section_tokens():
|
||||||
|
with pytest.raises(AuditaConfigError, match="AUDITA_MIN_SECTION_TOKENS"):
|
||||||
|
AuditaConfig.from_sources(
|
||||||
|
env={
|
||||||
|
"AUDITA_MIN_SECTION_TOKENS": "9000",
|
||||||
|
"AUDITA_MAX_SECTION_TOKENS": "8000",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import json
|
import json
|
||||||
|
import threading
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
@@ -12,25 +13,42 @@ from audita.pipeline import process_transcript, process_transcript_result
|
|||||||
|
|
||||||
class FakeStructuredLLMClient:
|
class FakeStructuredLLMClient:
|
||||||
def __init__(self, responses):
|
def __init__(self, responses):
|
||||||
self._responses = list(responses)
|
self._responses = responses
|
||||||
|
self._lock = threading.Lock()
|
||||||
self.calls = []
|
self.calls = []
|
||||||
|
|
||||||
def run_structured(self, *, stage_name, messages, response_model, config):
|
def run_structured(self, *, stage_name, messages, response_model, config):
|
||||||
self.calls.append(
|
with self._lock:
|
||||||
{
|
self.calls.append(
|
||||||
"stage_name": stage_name,
|
{
|
||||||
"messages": list(messages),
|
"stage_name": stage_name,
|
||||||
"response_model": response_model,
|
"messages": list(messages),
|
||||||
}
|
"response_model": response_model,
|
||||||
)
|
}
|
||||||
if not self._responses:
|
)
|
||||||
raise AuditaLLMError("FakeStructuredLLMClient received more calls than expected.")
|
response = _pop_llm_response(self._responses, stage_name)
|
||||||
response = self._responses.pop(0)
|
|
||||||
if isinstance(response, Exception):
|
if isinstance(response, Exception):
|
||||||
raise response
|
raise response
|
||||||
return response_model.model_validate(response)
|
return response_model.model_validate(response)
|
||||||
|
|
||||||
|
|
||||||
|
def _pop_llm_response(responses, stage_name):
|
||||||
|
if isinstance(responses, dict):
|
||||||
|
if stage_name not in responses:
|
||||||
|
raise AuditaLLMError(f"FakeStructuredLLMClient received unexpected stage_name: {stage_name}")
|
||||||
|
payloads = responses[stage_name]
|
||||||
|
if isinstance(payloads, list):
|
||||||
|
if not payloads:
|
||||||
|
raise AuditaLLMError(f"FakeStructuredLLMClient received too many calls for stage_name: {stage_name}")
|
||||||
|
return payloads.pop(0)
|
||||||
|
payload = payloads
|
||||||
|
del responses[stage_name]
|
||||||
|
return payload
|
||||||
|
if not responses:
|
||||||
|
raise AuditaLLMError("FakeStructuredLLMClient received more calls than expected.")
|
||||||
|
return responses.pop(0)
|
||||||
|
|
||||||
|
|
||||||
def _glossary():
|
def _glossary():
|
||||||
return parse_glossary_yaml(
|
return parse_glossary_yaml(
|
||||||
"""
|
"""
|
||||||
|
|||||||
Reference in New Issue
Block a user