Skip to content

Conditional

ConditionalUnionExtractor for multi-pass extraction with chat history passed between passes.

Classes:

Name Description
ConditionalUnionExtractor

Extends UnionExtractor by feeding each pass's response into the chat history for subsequent passes.

ConditionalUnionExtractor(overrides, aggregator, return_as_list=None, **kwargs)

Bases: UnionExtractor

Extractor that repeats extraction multiple times with history and aggregates results per key. This extractor calls the base extraction function multiple times (for each entry in overrides) on the same input text, passing the history of previous messages to each subsequent call.

See UnionExtractor for accepted parameters and details about the aggregation logic.

Source code in src/kibad_llm/extractors/union.py
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
def __init__(
    self,
    overrides: list[dict] | dict[str, dict],
    aggregator: Aggregator,
    return_as_list: list[str] | None = None,
    **kwargs,
):
    if len(overrides) < 1:
        raise ValueError("overrides must contain at least one set of parameters")
    if isinstance(overrides, list):
        overrides = {str(i): override for i, override in enumerate(overrides)}
    self.overrides = overrides
    self.aggregator = aggregator
    self.return_as_list = return_as_list or []
    self.default_kwargs = kwargs

__call__(*args, **kwargs)

Process singular text in multiple passes with chat history.

Parameters:

Name Type Description Default
*args Any

Are forwarded unchanged: extract_from_text_lenient

()

Other Parameters:

Name Type Description
* Any

Returns:

Type Description
dict[str, Any]

Dict with the key structured that holds the aggregated structured outputs.

dict[str, Any]

Additionally there can be lists for fields at the keys "{field}_list".

Source code in src/kibad_llm/extractors/conditional.py
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
def __call__(self, *args, **kwargs) -> dict[str, Any]:
    """Process singular text in multiple passes with chat history.

    Args:
        *args (Any): Are forwarded unchanged:
            [`extract_from_text_lenient`][kibad_llm.extractors.base.extract_from_text_lenient]

    Keyword Args:
        * (Any): Refer to [`extract_from_text_lenient`][kibad_llm.extractors.base.extract_from_text_lenient]

    Returns:
        Dict with the key `structured` that holds the aggregated structured outputs.
        Additionally there can be lists for fields at the keys `"{field}_list"`.
    """
    combined_kwargs = {**self.default_kwargs, **kwargs}
    results = []
    history: list[SimpleChatMessage] = []
    for override_name, override_params in self.overrides.items():
        # adjust kwargs:
        # 1) to return formatted messages for history
        current_kwargs = {
            **combined_kwargs,
            **override_params,
            "return_messages_formatted": True,
            "truncate_user_message_formatted": None,
        }
        # 2) if history exists, pass it and disable system message
        if len(history) > 0:
            current_kwargs["prompt_template"]["system_message"] = None
            current_kwargs["history"] = history

        current_result = extract_from_text_lenient(*args, **current_kwargs)

        # collect messages for history
        for role_str, content in current_result["messages_formatted"].items():
            role = MessageRole(role_str)
            history.append(SimpleChatMessage(role=role, content=content))
        # append assistant response or error to history
        history.append(
            SimpleChatMessage(
                role=MessageRole.ASSISTANT,
                content=current_result["response_content"] or current_result["error"],
            )
        )

        results.append(current_result)

    structured_outputs = [v.get("structured", None) for v in results]
    aggregated_structured = self.aggregator(structured_outputs)

    result: dict[str, Any] = {
        "structured": aggregated_structured,
    }
    for field in self.return_as_list:
        result[f"{field}_list"] = [v.get(field, None) for v in results]
    return result