diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 5e9680d..626952b 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -24,13 +24,11 @@ jobs: - name: Install dependencies run: | - cd v2 - uv pip install -e ".[dev]" + uv sync --extra dev - name: Run unit tests run: | - cd v2 - pytest tests/ -v -m "not integration" --tb=short + uv run pytest tests/ -v -m "not integration" --tb=short docker-build: runs-on: ubuntu-latest @@ -42,12 +40,11 @@ jobs: - name: Build Docker image run: | - cd v2 - docker build -f docker/codex-runtime.Dockerfile -t agentic-datagen-codex:v2 . + docker build -f docker/codex-runtime.Dockerfile -t teich-codex:ci . - name: Test Codex CLI in container run: | - docker run --rm agentic-datagen-codex:v2 codex --version + docker run --rm teich-codex:ci codex --version integration-tests: runs-on: ubuntu-latest @@ -70,12 +67,10 @@ jobs: - name: Install dependencies run: | - cd v2 - uv pip install -e ".[dev]" + uv sync --extra dev - name: Run integration tests env: OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }} run: | - cd v2 - pytest tests/test_integration.py -v --tb=short + uv run pytest tests/test_integration.py -v --tb=short diff --git a/pyproject.toml b/pyproject.toml index cf3598a..6ed1382 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,14 +20,13 @@ dev = ["pytest>=8.0", "pytest-asyncio>=0.23", "ruff>=0.4", "respx>=0.22"] [project.scripts] teich = "teich.cli:main" -agentic-datagen = "teich.cli:main" [build-system] requires = ["hatchling"] build-backend = "hatchling.build" [tool.hatch.build.targets.wheel] -packages = ["src/agentic_datagen", "src/teich"] +packages = ["src/teich"] [tool.pytest.ini_options] testpaths = ["tests"] diff --git a/src/agentic_datagen/__init__.py b/src/agentic_datagen/__init__.py deleted file mode 100644 index 285ffb9..0000000 --- a/src/agentic_datagen/__init__.py +++ /dev/null @@ -1,18 +0,0 @@ -"""Agentic Datagen v2 - Generate training data from Codex and Pi traces.""" - -__version__ = "0.1.1a3" - -from .config import Config, load_config -from .converter import TrainingExample, convert_trace_to_training_example, convert_traces_to_training_data -from .formatter import format_and_mask -from .loader import load_traces - -__all__ = [ - "Config", - "TrainingExample", - "convert_trace_to_training_example", - "convert_traces_to_training_data", - "format_and_mask", - "load_traces", - "load_config", -] diff --git a/src/agentic_datagen/__main__.py b/src/agentic_datagen/__main__.py deleted file mode 100644 index 9ae637f..0000000 --- a/src/agentic_datagen/__main__.py +++ /dev/null @@ -1,4 +0,0 @@ -from .cli import main - -if __name__ == "__main__": - main() diff --git a/src/teich/__init__.py b/src/teich/__init__.py index 53d822d..aede90f 100644 --- a/src/teich/__init__.py +++ b/src/teich/__init__.py @@ -1,29 +1,11 @@ -from __future__ import annotations +"""Teich - generate training data from Codex and Pi traces.""" -import sys -from importlib import import_module +from .config import Config, load_config +from .converter import TrainingExample, convert_trace_to_training_example, convert_traces_to_training_data +from .formatter import format_and_mask +from .loader import load_traces -from agentic_datagen import ( - Config, - TrainingExample, - __version__, - convert_trace_to_training_example, - convert_traces_to_training_data, - format_and_mask, - load_config, - load_traces, -) - -for _module_name in ( - "cli", - "config", - "converter", - "formatter", - "loader", - "runner", - "trace_readme", -): - sys.modules[f"{__name__}.{_module_name}"] = import_module(f"agentic_datagen.{_module_name}") +__version__ = "0.1.1a3" __all__ = [ "Config", diff --git a/src/agentic_datagen/cli.py b/src/teich/cli.py similarity index 100% rename from src/agentic_datagen/cli.py rename to src/teich/cli.py diff --git a/src/agentic_datagen/config.py b/src/teich/config.py similarity index 82% rename from src/agentic_datagen/config.py rename to src/teich/config.py index f951b24..1906366 100644 --- a/src/agentic_datagen/config.py +++ b/src/teich/config.py @@ -2,7 +2,6 @@ from __future__ import annotations -import csv import os from pathlib import Path import re @@ -10,6 +9,8 @@ import yaml from pydantic import BaseModel, Field, field_validator, model_validator +from .utils.prompts import load_prompt_rows + GITHUB_REPO_ID_PATTERN = re.compile(r"^[A-Za-z0-9_.-]+/[A-Za-z0-9_.-]+$") @@ -225,7 +226,14 @@ def get_prompt_inputs(self) -> list[PromptInput]: """Get structured prompt inputs from config and prompts_file.""" prompt_inputs = [PromptInput(prompt=prompt) for prompt in self.prompts] if self.prompts_file: - prompt_inputs.extend(self._load_prompt_inputs_from_file(self.prompts_file)) + prompt_inputs.extend( + PromptInput( + image=row.get("image"), + github_repo=row.get("github_repo"), + prompt=row.get("prompt") or "", + ) + for row in load_prompt_rows(self.prompts_file) + ) for prompt_input in prompt_inputs: if prompt_input.image is not None: raise ValueError( @@ -233,49 +241,6 @@ def get_prompt_inputs(self) -> list[PromptInput]: ) return prompt_inputs - @staticmethod - def _load_prompt_inputs_from_file(path: Path) -> list[PromptInput]: - if path.suffix.lower() == ".csv": - return Config._load_prompt_inputs_from_csv(path) - return Config._load_prompt_inputs_from_text(path) - - @staticmethod - def _load_prompt_inputs_from_text(path: Path) -> list[PromptInput]: - with path.open("r", encoding="utf-8") as handle: - return [ - PromptInput(prompt=line.strip()) - for line in handle - if line.strip() and not line.startswith("#") - ] - - @staticmethod - def _load_prompt_inputs_from_csv(path: Path) -> list[PromptInput]: - with path.open("r", encoding="utf-8", newline="") as handle: - reader = csv.DictReader(handle) - fieldnames = [name.strip().lower() for name in reader.fieldnames or [] if isinstance(name, str)] - if "prompt" not in fieldnames: - raise ValueError("Prompt CSV must include a 'prompt' column") - prompt_inputs: list[PromptInput] = [] - for row in reader: - normalized_row = { - key.strip().lower(): value - for key, value in row.items() - if isinstance(key, str) - } - if not any( - isinstance(value, str) and value.strip() - for value in normalized_row.values() - ): - continue - prompt_inputs.append( - PromptInput( - image=normalized_row.get("image"), - github_repo=normalized_row.get("github_repo"), - prompt=normalized_row.get("prompt") or "", - ) - ) - return prompt_inputs - def load_config(path: Path) -> Config: """Load configuration from YAML file. diff --git a/src/agentic_datagen/converter.py b/src/teich/converter.py similarity index 60% rename from src/agentic_datagen/converter.py rename to src/teich/converter.py index f647f6a..1257a0d 100644 --- a/src/agentic_datagen/converter.py +++ b/src/teich/converter.py @@ -5,6 +5,17 @@ from pathlib import Path from typing import Any +from .utils.schema import infer_tool_parameters_schema +from .utils.trace import ( + first_text_block, + has_message, + is_tool_not_found_result, + parse_function_arguments, + parse_tool_descriptions, + pi_reasoning_text, + reasoning_summary, +) + @dataclass(slots=True) class TrainingExample: @@ -23,245 +34,6 @@ def to_dict(self) -> dict[str, Any]: } -def _first_text_block(content_blocks: Any) -> str: - if isinstance(content_blocks, str): - return content_blocks.strip() - if not isinstance(content_blocks, list): - return "" - parts: list[str] = [] - for block in content_blocks: - if not isinstance(block, dict): - continue - block_type = block.get("type") - if block_type in {"input_text", "output_text", "text"}: - text = block.get("text") - if isinstance(text, str) and text: - parts.append(text) - return "\n".join(parts).strip() - - -def _has_same_system_message(messages: list[dict[str, Any]], content: str) -> bool: - return any( - message.get("role") == "system" and message.get("content") == content - for message in messages - ) - - -def _pi_reasoning_content(content_blocks: Any) -> str | None: - if not isinstance(content_blocks, list): - return None - parts: list[str] = [] - for block in content_blocks: - if not isinstance(block, dict): - continue - if block.get("type") != "thinking": - continue - thinking = block.get("thinking") - if isinstance(thinking, str) and thinking.strip(): - parts.append(thinking.strip()) - result = "\n\n".join(parts).strip() - return result or None - - -def _tool_result_content_text(payload: dict[str, Any]) -> str: - return _first_text_block(payload.get("content")) - - -def _is_tool_not_found_result(tool_name: str | None, payload: dict[str, Any]) -> bool: - content = _tool_result_content_text(payload).strip() - if tool_name: - return content == f"Tool {tool_name} not found" - return content == "Tool not found" - - -def _reasoning_summary(payload: dict[str, Any]) -> str | None: - summary = payload.get("summary") - parts: list[str] = [] - if isinstance(summary, list): - for item in summary: - if not isinstance(item, dict): - continue - text = item.get("text") - if isinstance(text, str) and text.strip(): - parts.append(text.strip()) - result = "\n\n".join(parts).strip() - if result: - return result - - content = payload.get("content") - if not isinstance(content, list): - return None - for item in content: - if not isinstance(item, dict): - continue - if item.get("type") != "reasoning_text": - continue - text = item.get("text") - if isinstance(text, str) and text.strip(): - parts.append(text.strip()) - result = "\n\n".join(parts).strip() - return result or None - - -def _normalize_json_like_value(value: Any) -> Any: - if isinstance(value, dict): - return {key: _normalize_json_like_value(item) for key, item in value.items()} - if isinstance(value, list): - return [_normalize_json_like_value(item) for item in value] - if not isinstance(value, str): - return value - stripped = value.strip() - if not stripped or stripped[0] not in "[{": - return value - try: - parsed = json.loads(stripped) - except json.JSONDecodeError: - return value - return _normalize_json_like_value(parsed) - - -def _parse_function_arguments(arguments: Any) -> Any: - if not isinstance(arguments, str): - return _normalize_json_like_value(arguments) if arguments is not None else {} - stripped = arguments.strip() - if not stripped: - return {} - try: - return _normalize_json_like_value(json.loads(stripped)) - except json.JSONDecodeError: - return arguments - - -def _schema_identity(schema: dict[str, Any]) -> str: - return json.dumps(schema, sort_keys=True, ensure_ascii=False) - - -def _infer_schema_from_value(value: Any) -> dict[str, Any]: - if value is None: - return {"type": "null"} - if isinstance(value, bool): - return {"type": "boolean"} - if isinstance(value, int): - return {"type": "integer"} - if isinstance(value, float): - return {"type": "number"} - if isinstance(value, str): - return {"type": "string"} - if isinstance(value, list): - item_schemas = [_infer_schema_from_value(item) for item in value] - schema: dict[str, Any] = {"type": "array"} - if item_schemas: - schema["items"] = _merge_schemas(item_schemas) - return schema - if isinstance(value, dict): - return _infer_tool_parameters_schema([value]) - return {} - - -def _merge_object_schemas(schemas: list[dict[str, Any]]) -> dict[str, Any]: - properties_by_name: dict[str, list[dict[str, Any]]] = {} - required_sets: list[set[str]] = [] - additional_properties = False - for schema in schemas: - properties = schema.get("properties") - if isinstance(properties, dict): - for name, value in properties.items(): - if isinstance(value, dict): - properties_by_name.setdefault(name, []).append(value) - required = schema.get("required") - if isinstance(required, list): - required_sets.append({item for item in required if isinstance(item, str)}) - else: - required_sets.append(set()) - if schema.get("additionalProperties", True) is not False: - additional_properties = True - merged: dict[str, Any] = { - "type": "object", - "properties": { - name: _merge_schemas(property_schemas) - for name, property_schemas in sorted(properties_by_name.items()) - }, - "additionalProperties": additional_properties, - } - if required_sets: - required = sorted(set.intersection(*required_sets)) - if required: - merged["required"] = required - return merged - - -def _merge_schemas(schemas: list[dict[str, Any]]) -> dict[str, Any]: - unique: list[dict[str, Any]] = [] - seen: set[str] = set() - for schema in schemas: - if not schema: - continue - identity = _schema_identity(schema) - if identity in seen: - continue - seen.add(identity) - unique.append(schema) - if not unique: - return {} - if len(unique) == 1: - return unique[0] - schema_types = {schema.get("type") for schema in unique if isinstance(schema.get("type"), str)} - if schema_types == {"object"}: - return _merge_object_schemas(unique) - if schema_types == {"array"}: - item_schemas = [schema.get("items") for schema in unique if isinstance(schema.get("items"), dict)] - merged: dict[str, Any] = {"type": "array"} - if item_schemas: - merged["items"] = _merge_schemas(item_schemas) - return merged - return {"anyOf": unique} - - -def _infer_tool_parameters_schema(argument_samples: list[Any]) -> dict[str, Any]: - dict_samples = [sample for sample in argument_samples if isinstance(sample, dict)] - if not dict_samples: - return {"type": "object", "properties": {}, "additionalProperties": True} - properties: dict[str, dict[str, Any]] = {} - all_keys = sorted({key for sample in dict_samples for key in sample}) - for key in all_keys: - observed = [_infer_schema_from_value(sample[key]) for sample in dict_samples if key in sample] - properties[key] = _merge_schemas(observed) - required = sorted(set.intersection(*(set(sample.keys()) for sample in dict_samples))) if dict_samples else [] - schema: dict[str, Any] = { - "type": "object", - "properties": properties, - "additionalProperties": True, - } - if required: - schema["required"] = required - return schema - - -def _parse_tool_descriptions_from_text(text: str) -> dict[str, str]: - descriptions: dict[str, str] = {} - in_section = False - for raw_line in text.splitlines(): - line = raw_line.strip() - if not in_section: - if line == "Available tools:": - in_section = True - continue - if not line: - if descriptions: - break - continue - if not line.startswith("- "): - if descriptions: - break - continue - name, separator, description = line[2:].partition(":") - tool_name = name.strip() - tool_description = description.strip() - if separator and tool_name and tool_description: - descriptions[tool_name] = tool_description - return descriptions - - def _normalize_role(role: str) -> str: if role == "developer": return "system" @@ -334,9 +106,9 @@ def _convert_codex_trace_to_training_example( base_instructions = payload.get("base_instructions") if isinstance(base_instructions, dict): text = base_instructions.get("text") - if isinstance(text, str) and text.strip() and not _has_same_system_message(messages, text): + if isinstance(text, str) and text.strip() and not has_message(messages, role="system", content=text): messages.append({"role": "system", "content": text}) - tool_descriptions.update(_parse_tool_descriptions_from_text(text)) + tool_descriptions.update(parse_tool_descriptions(text)) continue if event_type == "turn_context" and isinstance(payload, dict): turn_contexts.append(payload) @@ -346,7 +118,7 @@ def _convert_codex_trace_to_training_example( payload_type = payload.get("type") if payload_type == "reasoning": - pending_reasoning = _reasoning_summary(payload) + pending_reasoning = reasoning_summary(payload) continue if payload_type == "message": @@ -354,7 +126,7 @@ def _convert_codex_trace_to_training_example( if not isinstance(role, str): continue normalized_role = _normalize_role(role) - content = _first_text_block(payload.get("content")) + content = first_text_block(payload.get("content")) if normalized_role == "user" and content and not prompt: prompt = content message: dict[str, Any] = { @@ -374,7 +146,7 @@ def _convert_codex_trace_to_training_example( continue tool_names.add(name) tool_call_names[call_id] = name - arguments = _parse_function_arguments(payload.get("arguments")) + arguments = parse_function_arguments(payload.get("arguments")) tool_argument_samples.setdefault(name, []).append(arguments) tool_call = { "id": call_id, @@ -428,7 +200,7 @@ def _convert_codex_trace_to_training_example( if name in tool_descriptions and "description" not in schema: schema["description"] = tool_descriptions[name] if "parameters" not in schema: - schema["parameters"] = _infer_tool_parameters_schema(tool_argument_samples.get(name, [])) + schema["parameters"] = infer_tool_parameters_schema(tool_argument_samples.get(name, [])) tools.append(_build_tool_entry(name, schema)) if not prompt: prompt = next( @@ -483,7 +255,7 @@ def _convert_pi_trace_to_training_example( if not isinstance(tool_call_id, str) or not tool_call_id: continue tool_name = payload.get("toolName") if isinstance(payload.get("toolName"), str) else None - if _is_tool_not_found_result(tool_name, payload): + if is_tool_not_found_result(tool_name, payload): invalid_tool_call_ids.add(tool_call_id) for event in events: @@ -527,7 +299,7 @@ def _convert_pi_trace_to_training_example( "role": "tool", "tool_call_id": tool_call_id, "name": tool_name or "unknown_tool", - "content": _first_text_block(payload.get("content")), + "content": first_text_block(payload.get("content")), } if payload.get("isError") is True: tool_message["is_error"] = True @@ -536,10 +308,10 @@ def _convert_pi_trace_to_training_example( normalized_role = _normalize_role(role) content_blocks = payload.get("content") - content = _first_text_block(content_blocks) + content = first_text_block(content_blocks) if role == "developer" and content: - tool_descriptions.update(_parse_tool_descriptions_from_text(content)) + tool_descriptions.update(parse_tool_descriptions(content)) if normalized_role == "user": if content and not prompt: @@ -552,7 +324,7 @@ def _convert_pi_trace_to_training_example( "content": content, } if normalized_role == "assistant": - reasoning_content = _pi_reasoning_content(content_blocks) + reasoning_content = pi_reasoning_text(content_blocks) if reasoning_content: message["reasoning_content"] = reasoning_content tool_calls: list[dict[str, Any]] = [] @@ -569,7 +341,7 @@ def _convert_pi_trace_to_training_example( if not tool_call_id or not tool_name or tool_call_id in invalid_tool_call_ids: continue tool_names.add(tool_name) - arguments = _parse_function_arguments(block.get("arguments")) + arguments = parse_function_arguments(block.get("arguments")) tool_argument_samples.setdefault(tool_name, []).append(arguments) tool_calls.append( { @@ -594,7 +366,7 @@ def _convert_pi_trace_to_training_example( name, { **({"description": tool_descriptions[name]} if name in tool_descriptions else {}), - "parameters": _infer_tool_parameters_schema(tool_argument_samples.get(name, [])), + "parameters": infer_tool_parameters_schema(tool_argument_samples.get(name, [])), }, ) for name in sorted(tool_names) diff --git a/src/agentic_datagen/formatter.py b/src/teich/formatter.py similarity index 100% rename from src/agentic_datagen/formatter.py rename to src/teich/formatter.py diff --git a/src/agentic_datagen/loader.py b/src/teich/loader.py similarity index 100% rename from src/agentic_datagen/loader.py rename to src/teich/loader.py diff --git a/src/agentic_datagen/runner.py b/src/teich/runner.py similarity index 99% rename from src/agentic_datagen/runner.py rename to src/teich/runner.py index 248125e..ff841a7 100644 --- a/src/agentic_datagen/runner.py +++ b/src/teich/runner.py @@ -6,7 +6,6 @@ from concurrent.futures import ThreadPoolExecutor, as_completed from dataclasses import dataclass import json -import os import re import shlex import shutil diff --git a/src/agentic_datagen/trace_readme.py b/src/teich/trace_readme.py similarity index 72% rename from src/agentic_datagen/trace_readme.py rename to src/teich/trace_readme.py index 8d5f876..fb111a2 100644 --- a/src/agentic_datagen/trace_readme.py +++ b/src/teich/trace_readme.py @@ -5,54 +5,7 @@ from typing import Any, Iterable from .converter import convert_trace_to_training_example - - -def _merge_tool_parameters(schemas: list[dict[str, Any]]) -> dict[str, Any]: - object_schemas = [schema for schema in schemas if isinstance(schema, dict) and schema] - if not object_schemas: - return {"type": "object", "properties": {}, "additionalProperties": True} - if len(object_schemas) == 1: - return object_schemas[0] - properties: dict[str, list[dict[str, Any]]] = {} - required_sets: list[set[str]] = [] - additional_properties = False - for schema in object_schemas: - schema_properties = schema.get("properties") - if isinstance(schema_properties, dict): - for key, value in schema_properties.items(): - if isinstance(value, dict): - properties.setdefault(key, []).append(value) - required = schema.get("required") - if isinstance(required, list): - required_sets.append({item for item in required if isinstance(item, str)}) - else: - required_sets.append(set()) - if schema.get("additionalProperties", True) is not False: - additional_properties = True - merged_properties: dict[str, dict[str, Any]] = {} - for key, values in sorted(properties.items()): - unique_values: list[dict[str, Any]] = [] - seen: set[str] = set() - for value in values: - identity = json.dumps(value, sort_keys=True, ensure_ascii=False) - if identity in seen: - continue - seen.add(identity) - unique_values.append(value) - if len(unique_values) == 1: - merged_properties[key] = unique_values[0] - else: - merged_properties[key] = {"anyOf": unique_values} - merged: dict[str, Any] = { - "type": "object", - "properties": merged_properties, - "additionalProperties": additional_properties, - } - if required_sets: - required = sorted(set.intersection(*required_sets)) - if required: - merged["required"] = required - return merged +from .utils.schema import merge_schemas def _dataset_tools(trace_files: Iterable[Path]) -> list[dict[str, Any]]: @@ -83,7 +36,7 @@ def _dataset_tools(trace_files: Iterable[Path]) -> list[dict[str, Any]]: existing_schema = merged_function.get("parameters") schema_list = [existing_schema] if isinstance(existing_schema, dict) else [] schema_list.append(schema) - merged_function["parameters"] = _merge_tool_parameters(schema_list) + merged_function["parameters"] = merge_schemas(schema_list) return [merged_by_name[name] for name in sorted(merged_by_name)] diff --git a/src/teich/utils/__init__.py b/src/teich/utils/__init__.py new file mode 100644 index 0000000..8f37254 --- /dev/null +++ b/src/teich/utils/__init__.py @@ -0,0 +1 @@ +"""Internal utility helpers for teich.""" diff --git a/src/teich/utils/prompts.py b/src/teich/utils/prompts.py new file mode 100644 index 0000000..d2ff7b2 --- /dev/null +++ b/src/teich/utils/prompts.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +import csv +from pathlib import Path + + +PromptRow = dict[str, str | None] + + +def load_prompt_rows(path: Path) -> list[PromptRow]: + if path.suffix.lower() == ".csv": + return _load_prompt_rows_from_csv(path) + return _load_prompt_rows_from_text(path) + + +def _load_prompt_rows_from_text(path: Path) -> list[PromptRow]: + with path.open("r", encoding="utf-8") as handle: + return [ + {"prompt": line.strip()} + for line in handle + if line.strip() and not line.startswith("#") + ] + + +def _load_prompt_rows_from_csv(path: Path) -> list[PromptRow]: + with path.open("r", encoding="utf-8", newline="") as handle: + reader = csv.DictReader(handle) + fieldnames = [name.strip().lower() for name in reader.fieldnames or [] if isinstance(name, str)] + if "prompt" not in fieldnames: + raise ValueError("Prompt CSV must include a 'prompt' column") + + rows: list[PromptRow] = [] + for row in reader: + normalized_row: PromptRow = { + key.strip().lower(): value + for key, value in row.items() + if isinstance(key, str) + } + if not any(isinstance(value, str) and value.strip() for value in normalized_row.values()): + continue + rows.append(normalized_row) + return rows diff --git a/src/teich/utils/schema.py b/src/teich/utils/schema.py new file mode 100644 index 0000000..5b83129 --- /dev/null +++ b/src/teich/utils/schema.py @@ -0,0 +1,119 @@ +from __future__ import annotations + +import json +from typing import Any + + +def empty_object_schema() -> dict[str, Any]: + return {"type": "object", "properties": {}, "additionalProperties": True} + + +def schema_identity(schema: dict[str, Any]) -> str: + return json.dumps(schema, sort_keys=True, ensure_ascii=False) + + +def merge_object_schemas(schemas: list[dict[str, Any]]) -> dict[str, Any]: + properties_by_name: dict[str, list[dict[str, Any]]] = {} + required_sets: list[set[str]] = [] + additional_properties = False + + for schema in schemas: + properties = schema.get("properties") + if isinstance(properties, dict): + for name, value in properties.items(): + if isinstance(value, dict): + properties_by_name.setdefault(name, []).append(value) + required = schema.get("required") + if isinstance(required, list): + required_sets.append({item for item in required if isinstance(item, str)}) + else: + required_sets.append(set()) + if schema.get("additionalProperties", True) is not False: + additional_properties = True + + merged: dict[str, Any] = { + "type": "object", + "properties": { + name: merge_schemas(property_schemas) + for name, property_schemas in sorted(properties_by_name.items()) + }, + "additionalProperties": additional_properties, + } + if required_sets: + required = sorted(set.intersection(*required_sets)) + if required: + merged["required"] = required + return merged + + +def merge_schemas(schemas: list[dict[str, Any]]) -> dict[str, Any]: + unique: list[dict[str, Any]] = [] + seen: set[str] = set() + for schema in schemas: + if not schema: + continue + identity = schema_identity(schema) + if identity in seen: + continue + seen.add(identity) + unique.append(schema) + + if not unique: + return {} + if len(unique) == 1: + return unique[0] + + schema_types = {schema.get("type") for schema in unique if isinstance(schema.get("type"), str)} + if schema_types == {"object"}: + return merge_object_schemas(unique) + if schema_types == {"array"}: + item_schemas = [schema.get("items") for schema in unique if isinstance(schema.get("items"), dict)] + merged: dict[str, Any] = {"type": "array"} + if item_schemas: + merged["items"] = merge_schemas(item_schemas) + return merged + return {"anyOf": unique} + + +def infer_schema_from_value(value: Any) -> dict[str, Any]: + if value is None: + return {"type": "null"} + if isinstance(value, bool): + return {"type": "boolean"} + if isinstance(value, int): + return {"type": "integer"} + if isinstance(value, float): + return {"type": "number"} + if isinstance(value, str): + return {"type": "string"} + if isinstance(value, list): + item_schemas = [infer_schema_from_value(item) for item in value] + schema: dict[str, Any] = {"type": "array"} + if item_schemas: + schema["items"] = merge_schemas(item_schemas) + return schema + if isinstance(value, dict): + return infer_tool_parameters_schema([value]) + return {} + + +def infer_tool_parameters_schema(argument_samples: list[Any]) -> dict[str, Any]: + dict_samples = [sample for sample in argument_samples if isinstance(sample, dict)] + if not dict_samples: + return empty_object_schema() + + properties: dict[str, dict[str, Any]] = {} + all_keys = sorted({key for sample in dict_samples for key in sample}) + for key in all_keys: + observed = [infer_schema_from_value(sample[key]) for sample in dict_samples if key in sample] + properties[key] = merge_schemas(observed) + + required = sorted(set.intersection(*(set(sample.keys()) for sample in dict_samples))) if dict_samples else [] + schema: dict[str, Any] = { + "type": "object", + "properties": properties, + "additionalProperties": True, + } + if required: + schema["required"] = required + return schema diff --git a/src/teich/utils/trace.py b/src/teich/utils/trace.py new file mode 100644 index 0000000..b06ccd0 --- /dev/null +++ b/src/teich/utils/trace.py @@ -0,0 +1,140 @@ +from __future__ import annotations + +import json +from typing import Any + + +def first_text_block(content_blocks: Any) -> str: + if isinstance(content_blocks, str): + return content_blocks.strip() + if not isinstance(content_blocks, list): + return "" + + parts: list[str] = [] + for block in content_blocks: + if not isinstance(block, dict): + continue + if block.get("type") not in {"input_text", "output_text", "text"}: + continue + text = block.get("text") + if isinstance(text, str) and text: + parts.append(text) + return "\n".join(parts).strip() + + +def has_message(messages: list[dict[str, Any]], *, role: str, content: str) -> bool: + return any(message.get("role") == role and message.get("content") == content for message in messages) + + +def pi_reasoning_text(content_blocks: Any) -> str | None: + if not isinstance(content_blocks, list): + return None + + parts: list[str] = [] + for block in content_blocks: + if not isinstance(block, dict) or block.get("type") != "thinking": + continue + thinking = block.get("thinking") + if isinstance(thinking, str) and thinking.strip(): + parts.append(thinking.strip()) + + result = "\n\n".join(parts).strip() + return result or None + + +def tool_result_content_text(payload: dict[str, Any]) -> str: + return first_text_block(payload.get("content")) + + +def is_tool_not_found_result(tool_name: str | None, payload: dict[str, Any]) -> bool: + content = tool_result_content_text(payload).strip() + if tool_name: + return content == f"Tool {tool_name} not found" + return content == "Tool not found" + + +def reasoning_summary(payload: dict[str, Any]) -> str | None: + summary = payload.get("summary") + parts: list[str] = [] + if isinstance(summary, list): + for item in summary: + if not isinstance(item, dict): + continue + text = item.get("text") + if isinstance(text, str) and text.strip(): + parts.append(text.strip()) + + result = "\n\n".join(parts).strip() + if result: + return result + + content = payload.get("content") + if not isinstance(content, list): + return None + for item in content: + if not isinstance(item, dict) or item.get("type") != "reasoning_text": + continue + text = item.get("text") + if isinstance(text, str) and text.strip(): + parts.append(text.strip()) + + result = "\n\n".join(parts).strip() + return result or None + + +def parse_tool_descriptions(text: str) -> dict[str, str]: + descriptions: dict[str, str] = {} + in_section = False + + for raw_line in text.splitlines(): + line = raw_line.strip() + if not in_section: + if line == "Available tools:": + in_section = True + continue + if not line: + if descriptions: + break + continue + if not line.startswith("- "): + if descriptions: + break + continue + + name, separator, description = line[2:].partition(":") + tool_name = name.strip() + tool_description = description.strip() + if separator and tool_name and tool_description: + descriptions[tool_name] = tool_description + return descriptions + + +def normalize_json_like_value(value: Any) -> Any: + if isinstance(value, dict): + return {key: normalize_json_like_value(item) for key, item in value.items()} + if isinstance(value, list): + return [normalize_json_like_value(item) for item in value] + if not isinstance(value, str): + return value + + stripped = value.strip() + if not stripped or stripped[0] not in "[{": + return value + try: + parsed = json.loads(stripped) + except json.JSONDecodeError: + return value + return normalize_json_like_value(parsed) + + +def parse_function_arguments(arguments: Any) -> Any: + if not isinstance(arguments, str): + return normalize_json_like_value(arguments) if arguments is not None else {} + + stripped = arguments.strip() + if not stripped: + return {} + try: + return normalize_json_like_value(json.loads(stripped)) + except json.JSONDecodeError: + return arguments diff --git a/tests/test_cli.py b/tests/test_cli.py index 04eb68f..f1dadbf 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -2,8 +2,6 @@ from pathlib import Path from unittest.mock import MagicMock, patch - -import pytest from rich.console import Console from typer.testing import CliRunner diff --git a/tests/test_config.py b/tests/test_config.py index 2fb9db3..6cdb2fc 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -1,11 +1,10 @@ """Tests for config module.""" -import tempfile from pathlib import Path import pytest -from teich.config import Config, MCPConfig, ModelConfig +from teich.config import Config, MCPConfig def test_default_config(): diff --git a/tests/test_integration.py b/tests/test_integration.py index 2014662..700af30 100644 --- a/tests/test_integration.py +++ b/tests/test_integration.py @@ -8,7 +8,6 @@ import json import os import subprocess -import tempfile from pathlib import Path from unittest.mock import patch, MagicMock @@ -202,8 +201,6 @@ class TestEndToEnd: def test_full_generation_workflow(self, tmp_path): """Test complete workflow: init -> generate -> verify output.""" - import shutil - # Setup project_dir = tmp_path / "test-project" project_dir.mkdir()