fix(cua): add Azure OpenAI support, fix max_tokens → max_completion_tokens, increase screencap timeout
This commit is contained in:
@@ -52,7 +52,7 @@ import time
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from openai import OpenAI
|
from openai import OpenAI, AzureOpenAI
|
||||||
except ImportError:
|
except ImportError:
|
||||||
sys.exit("openai package required: pip install openai")
|
sys.exit("openai package required: pip install openai")
|
||||||
|
|
||||||
@@ -79,7 +79,7 @@ def screenshot_b64() -> str:
|
|||||||
# Use exec-out for direct binary pipe (fastest, no intermediate file on device)
|
# Use exec-out for direct binary pipe (fastest, no intermediate file on device)
|
||||||
result = subprocess.run(
|
result = subprocess.run(
|
||||||
["adb", "exec-out", "screencap", "-p"],
|
["adb", "exec-out", "screencap", "-p"],
|
||||||
capture_output=True, timeout=10,
|
capture_output=True, timeout=30,
|
||||||
)
|
)
|
||||||
if result.returncode == 0 and len(result.stdout) > 100:
|
if result.returncode == 0 and len(result.stdout) > 100:
|
||||||
return base64.b64encode(result.stdout).decode()
|
return base64.b64encode(result.stdout).decode()
|
||||||
@@ -184,14 +184,14 @@ Rules:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
def call_llm(client: OpenAI, model: str, system: str, history: list) -> str:
|
def call_llm(client, model: str, system: str, history: list) -> str:
|
||||||
"""Call LLM via OpenAI-compatible API with retry on rate limit."""
|
"""Call LLM via OpenAI-compatible API with retry on rate limit."""
|
||||||
for attempt in range(3):
|
for attempt in range(3):
|
||||||
try:
|
try:
|
||||||
response = client.chat.completions.create(
|
response = client.chat.completions.create(
|
||||||
model=model,
|
model=model,
|
||||||
messages=[{"role": "system", "content": system}] + history,
|
messages=[{"role": "system", "content": system}] + history,
|
||||||
max_tokens=300,
|
max_completion_tokens=300,
|
||||||
temperature=0,
|
temperature=0,
|
||||||
)
|
)
|
||||||
return response.choices[0].message.content.strip()
|
return response.choices[0].message.content.strip()
|
||||||
@@ -205,7 +205,13 @@ def call_llm(client: OpenAI, model: str, system: str, history: list) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def make_client(model: str):
|
def make_client(model: str):
|
||||||
"""Create OpenAI client. Supports OPENAI_API_KEY, GEMINI_API_KEY, XAI_API_KEY."""
|
"""Create OpenAI client. Supports AZURE_OPENAI_*, OPENAI_API_KEY, GEMINI_API_KEY, XAI_API_KEY."""
|
||||||
|
if os.environ.get("AZURE_OPENAI_API_KEY"):
|
||||||
|
return AzureOpenAI(
|
||||||
|
api_key=os.environ["AZURE_OPENAI_API_KEY"],
|
||||||
|
azure_endpoint=os.environ["AZURE_OPENAI_ENDPOINT"],
|
||||||
|
api_version=os.environ.get("AZURE_OPENAI_API_VERSION", "2024-08-01-preview"),
|
||||||
|
), model
|
||||||
if os.environ.get("OPENAI_API_KEY"):
|
if os.environ.get("OPENAI_API_KEY"):
|
||||||
base = os.environ.get("OPENAI_BASE_URL")
|
base = os.environ.get("OPENAI_BASE_URL")
|
||||||
return OpenAI(base_url=base) if base else OpenAI(), model
|
return OpenAI(base_url=base) if base else OpenAI(), model
|
||||||
@@ -219,7 +225,7 @@ def make_client(model: str):
|
|||||||
api_key=os.environ["GEMINI_API_KEY"],
|
api_key=os.environ["GEMINI_API_KEY"],
|
||||||
base_url="https://generativelanguage.googleapis.com/v1beta/openai/",
|
base_url="https://generativelanguage.googleapis.com/v1beta/openai/",
|
||||||
), "gemini-2.0-flash"
|
), "gemini-2.0-flash"
|
||||||
sys.exit("Set OPENAI_API_KEY, XAI_API_KEY, or GEMINI_API_KEY")
|
sys.exit("Set AZURE_OPENAI_API_KEY, OPENAI_API_KEY, XAI_API_KEY, or GEMINI_API_KEY")
|
||||||
|
|
||||||
|
|
||||||
def run_cua(goal: str, max_steps: int = 30, model: str = "gpt-4o",
|
def run_cua(goal: str, max_steps: int = 30, model: str = "gpt-4o",
|
||||||
|
|||||||
Reference in New Issue
Block a user