Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 25 additions & 4 deletions src/outlines/types/dsl.py
Original file line number Diff line number Diff line change
Expand Up @@ -822,15 +822,36 @@ def python_types_to_terms(ptype: Any, recursion_depth: int = 0) -> Term:


def _get_enum_members(ptype: EnumMeta) -> List[Any]:
"""Get the members of an enum, including "enum of callables" members.

Regular Enum members (including ones whose value is a `functools.partial`)
are returned via normal iteration. An enum can also have a plain function
assigned directly as a member's value (e.g. ``c = some_function``);
Python's Enum machinery treats a bare function assigned in the class body
as a method rather than a member, so those don't show up through
iteration and we have to pick them out of `__dict__` instead.

The tricky part is telling that case apart from an actual instance method
defined with ``def`` directly in the class body (e.g. a `describe(self)`
helper) - both end up as plain `FunctionType` objects in `__dict__`. We
use `__qualname__` to distinguish them: a method defined inside the
class body is qualified with the class's own qualname (e.g.
``"SomeEnum.describe"``), while a function that was defined elsewhere and
merely assigned as a member's value keeps its original qualname (e.g.
``"some_function"``). Only the latter should be treated as a member.
"""
regular_members = [member.value for member in ptype] # type: ignore
function_members = []
for key, value in ptype.__dict__.items():
class_member_prefix = f"{ptype.__qualname__}."
function_members = [
value
for key, value in ptype.__dict__.items()
if (
isinstance(value, FunctionType)
and not (key.startswith('__') and key.endswith('__'))
and key != '_generate_next_value_' # Skip this specific method that causes issues
):
function_members.append(value)
and not value.__qualname__.startswith(class_member_prefix)
)
]
return regular_members + function_members


Expand Down
21 changes: 21 additions & 0 deletions tests/types/test_dsl.py
Original file line number Diff line number Diff line change
Expand Up @@ -663,6 +663,27 @@ class SomeEnum(Enum):
python_types_to_terms(bytes)


def test_dsl_enum_with_instance_method_is_not_a_choice():
"""A plain instance method defined in the enum's class body must not be
picked up as a generation choice. Both this method and a bare function
assigned as a member's value (`SomeEnum.c` above) show up as the same
`FunctionType` in `__dict__`, so the two have to be told apart by more
than just their type."""

class Color(Enum):
RED = "red"
GREEN = "green"
BLUE = "blue"

def describe(self) -> str:
return f"This color is {self.value}"

result = python_types_to_terms(Color)
assert isinstance(result, Alternatives)
assert len(result.terms) == 3
assert result.terms == [String("red"), String("green"), String("blue")]


def test_dsl_handle_literal():
literal = Literal["a", 1]
result = _handle_literal(get_args(literal))
Expand Down
Loading