ai-deception-openenv / inference.py
spidey121's picture
fix round 1
0f3eb4c
Raw
History Blame Contribute Delete
4.06 kB
import time
import os
import random
import traceback
import requests
import threading
from openai import OpenAI
from env.fake_server import run_server
from env.env import DeceptionEnv
from env.attacker import simulate_attack
# Import graders
from tasks.easy.grader import grade as easy_grade
from tasks.medium.grader import grade as medium_grade
from tasks.hard.grader import grade as hard_grade
# Reproducibility
random.seed(42)
# Wait for server to start
def wait_for_server():
for _ in range(15):
try:
r = requests.get("http://127.0.0.1:7860/status")
if r.status_code == 200:
print("Server ready")
return
except Exception:
pass
time.sleep(1)
print("Server not started, continuing...")
# Start fake server safely
try:
threading.Thread(target=run_server, daemon=True).start()
except Exception as e:
print("Server start error:", e)
# Wait for server
wait_for_server()
# Environment variables
API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")
MODEL_NAME = os.getenv("MODEL_NAME", "Qwen/Qwen2.5-7B-Instruct")
HF_TOKEN = os.getenv("HF_TOKEN")
# Do NOT crash if missing
if HF_TOKEN is None:
print("Warning: HF_TOKEN not set")
MAX_STEPS = 5
def choose_action(client, state):
try:
response = client.chat.completions.create(
model=MODEL_NAME,
messages=[
{
"role": "system",
"content": """You are a cybersecurity deception agent.
Choose one:
detect_attack, deploy_honeypot, fake_database, block_ip"""
},
{
"role": "user",
"content": f"Current state: {state}"
}
],
max_tokens=10,
temperature=0.4
)
action = response.choices[0].message.content.strip()
except Exception:
action = "detect_attack"
return action
def run_task(task_name):
try:
client = OpenAI(
base_url=API_BASE_URL,
api_key=HF_TOKEN
)
env = DeceptionEnv()
state = env.reset()
print(
f"[START] task={task_name} env=ai-deception-openenv model={MODEL_NAME}",
flush=True
)
rewards = []
done = False
for step in range(1, MAX_STEPS + 1):
try:
simulate_attack()
except Exception:
pass
try:
state = env.state()
except Exception:
pass
action = choose_action(client, state)
if random.random() < 0.3:
action = random.choice(env.action_space())
if action not in env.action_space():
action = "detect_attack"
state, reward, done, _ = env.step(action)
rewards.append(reward)
print(
f"[STEP] step={step} action={action} reward={reward:.2f} "
f"done={str(done).lower()} error=null",
flush=True
)
if done:
break
steps = len(rewards)
# grading
score = 0.0
if rewards:
if task_name == "easy":
score = easy_grade(rewards)
elif task_name == "medium":
score = medium_grade(rewards)
else:
score = hard_grade(rewards)
success = score >= 0.3
print(
f"[END] success={str(success).lower()} steps={steps} "
f"score={score:.2f} rewards={','.join(f'{r:.2f}' for r in rewards)}",
flush=True
)
except Exception:
traceback.print_exc()
print(
"[END] success=false steps=0 score=0.00 rewards=",
flush=True
)
# Run all tasks safely
try:
run_task("easy")
run_task("medium")
run_task("hard")
except Exception:
traceback.print_exc()
time.sleep(30)