negotiation-arena / models.py
Ashutosh Kumar
Negotiation Arena V2: full OpenEnv environment
21a8fc4
Raw
History Blame Contribute Delete
3.08 kB
"""Action / Observation / State types for the Negotiation Arena."""
from __future__ import annotations
from typing import Any, Dict, List, Optional
from pydantic import Field
from openenv.core import Action, Observation, State
class NegotiationAction(Action):
"""A move submitted by the (trainable) Vendor agent on its turn."""
action_type: str = Field(
...,
description=(
"One of: make_offer, accept_offer, counter_offer, concede, demand, "
"walk_away, ask_clarification"
),
)
episode_id: str = Field(default="", description="Episode id from /reset")
turn_number: int = Field(default=0, description="Current turn number")
agent_role: str = Field(default="vendor", description="The acting role")
proposed_terms: Optional[Dict[str, Any]] = Field(
default=None,
description="For make_offer / counter_offer: full or partial term map",
)
issue_name: Optional[str] = Field(
default=None, description="For concede / demand: which issue"
)
new_value: Optional[Any] = Field(
default=None, description="For concede / demand: new proposed value"
)
reasoning: Optional[str] = Field(
default=None, description="Free-text justification"
)
class NegotiationObservation(Observation):
"""What the Vendor sees after every turn."""
success: bool = Field(default=True)
message: str = Field(default="")
turn_number: int = Field(default=0)
max_turns: int = Field(default=30)
task_id: str = Field(default="")
episode_id: str = Field(default="")
contract_draft: str = Field(default="")
open_issues: List[Dict[str, Any]] = Field(default_factory=list)
opponent_last_action: Optional[Dict[str, Any]] = Field(default=None)
offer_history: List[Dict[str, Any]] = Field(default_factory=list)
turns_remaining: int = Field(default=0)
vendor_private_brief: Optional[Dict[str, Any]] = Field(default=None)
final_deal_struck: Optional[bool] = Field(default=None)
final_terms: Optional[Dict[str, Any]] = Field(default=None)
client_walked_away: Optional[bool] = Field(default=None)
grader_breakdown: Optional[Dict[str, Any]] = Field(default=None)
error_message: Optional[str] = Field(default=None)
class NegotiationState(State):
"""Full episode state returned by /state."""
task_id: str = Field(default="")
turn_number: int = Field(default=0)
max_turns: int = Field(default=30)
done: bool = Field(default=False)
deal_struck: bool = Field(default=False)
offer_history: List[Dict[str, Any]] = Field(default_factory=list)
current_terms: Dict[str, Any] = Field(default_factory=dict)
issues_status: Dict[str, str] = Field(default_factory=dict)
vendor_brief: Dict[str, Any] = Field(default_factory=dict)
client_brief: Dict[str, Any] = Field(default_factory=dict)
turn_rewards: List[float] = Field(default_factory=list)
total_reward: float = Field(default=0.0)
final_grader_breakdown: Optional[Dict[str, Any]] = Field(default=None)