vasanthfeb13 commited on
Commit
3f08542
·
verified ·
1 Parent(s): 6297f9f

Upload folder using huggingface_hub

Browse files
Files changed (3) hide show
  1. baseline.py +10 -6
  2. inference.py +27 -12
  3. server/validator_tests.log +3 -0
baseline.py CHANGED
@@ -41,8 +41,8 @@ SYSTEM_PROMPT = (
41
 
42
  @dataclass
43
  class BaselineConfig:
44
- provider: str = "blaxel"
45
- model: str = "sandbox-openai"
46
  fallback_provider: str = "cerebras"
47
  fallback_model: str = "llama3.1-8b"
48
  episodes_per_task: int = 1
@@ -251,7 +251,11 @@ def _resolve_api_key(provider: str) -> str:
251
  return os.getenv("CEREBRAS_API_KEY", "").strip()
252
  if provider == "blaxel":
253
  return os.getenv("BLAXEL_AUTHORIZATION", "").strip()
254
- return os.getenv("OPENAI_API_KEY", "").strip()
 
 
 
 
255
 
256
 
257
  def _resolve_model(provider: str, model: str | None) -> str:
@@ -309,7 +313,7 @@ def _build_client(provider: str, api_key: str, model: str) -> Any:
309
  return OpenAI(api_key=normalized_key, base_url=base_url, default_headers=default_headers)
310
  return OpenAI(api_key=normalized_key, base_url=base_url)
311
 
312
- openai_base_url = os.getenv("OPENAI_API_BASE_URL", "").strip()
313
  if openai_base_url:
314
  return OpenAI(api_key=normalized_key, base_url=openai_base_url)
315
  return OpenAI(api_key=normalized_key)
@@ -376,8 +380,8 @@ def run_baseline_with_fallback_sync(
376
 
377
  def main() -> None:
378
  parser = argparse.ArgumentParser(description="Run SOC triage baseline across all tasks.")
379
- parser.add_argument("--provider", default=os.getenv("AI_PROVIDER", "blaxel"))
380
- parser.add_argument("--model", default=os.getenv("AI_MODEL", "sandbox-openai"))
381
  parser.add_argument("--fallback-provider", default=os.getenv("AI_FALLBACK_PROVIDER", "cerebras"))
382
  parser.add_argument("--fallback-model", default=os.getenv("AI_FALLBACK_MODEL", "llama3.1-8b"))
383
  parser.add_argument("--episodes", type=int, default=1)
 
41
 
42
  @dataclass
43
  class BaselineConfig:
44
+ provider: str = "openai"
45
+ model: str = "gpt-4o-mini"
46
  fallback_provider: str = "cerebras"
47
  fallback_model: str = "llama3.1-8b"
48
  episodes_per_task: int = 1
 
251
  return os.getenv("CEREBRAS_API_KEY", "").strip()
252
  if provider == "blaxel":
253
  return os.getenv("BLAXEL_AUTHORIZATION", "").strip()
254
+ return (
255
+ os.getenv("OPENAI_API_KEY", "").strip()
256
+ or os.getenv("API_KEY", "").strip()
257
+ or os.getenv("HF_TOKEN", "").strip()
258
+ )
259
 
260
 
261
  def _resolve_model(provider: str, model: str | None) -> str:
 
313
  return OpenAI(api_key=normalized_key, base_url=base_url, default_headers=default_headers)
314
  return OpenAI(api_key=normalized_key, base_url=base_url)
315
 
316
+ openai_base_url = os.getenv("OPENAI_API_BASE_URL", "").strip() or os.getenv("API_BASE_URL", "").strip()
317
  if openai_base_url:
318
  return OpenAI(api_key=normalized_key, base_url=openai_base_url)
319
  return OpenAI(api_key=normalized_key)
 
380
 
381
  def main() -> None:
382
  parser = argparse.ArgumentParser(description="Run SOC triage baseline across all tasks.")
383
+ parser.add_argument("--provider", default=os.getenv("AI_PROVIDER", "openai"))
384
+ parser.add_argument("--model", default=os.getenv("AI_MODEL", "gpt-4o-mini"))
385
  parser.add_argument("--fallback-provider", default=os.getenv("AI_FALLBACK_PROVIDER", "cerebras"))
386
  parser.add_argument("--fallback-model", default=os.getenv("AI_FALLBACK_MODEL", "llama3.1-8b"))
387
  parser.add_argument("--episodes", type=int, default=1)
inference.py CHANGED
@@ -3,8 +3,9 @@ Inference script for SOC Triage environment.
3
 
4
  MANDATORY submission variables:
5
  - API_BASE_URL: OpenAI-compatible chat completions API base URL.
 
 
6
  - MODEL_NAME: model identifier for inference.
7
- - HF_TOKEN: API token.
8
 
9
  STDOUT FORMAT (mandatory):
10
  [START] task=<task_name> env=soc_triage_env model=<model_name>
@@ -26,10 +27,19 @@ from typing import Any, List, Optional
26
  # ---------------------------------------------------------------------------
27
  # Mandatory env vars
28
  # ---------------------------------------------------------------------------
29
- API_BASE_URL = os.getenv("API_BASE_URL", "https://vasanthfeb13-soc-triage-env.hf.space")
30
- MODEL_NAME = os.getenv("MODEL_NAME", "sandbox-openai")
31
- HF_TOKEN = os.getenv("HF_TOKEN")
 
 
 
 
32
  LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME")
 
 
 
 
 
33
 
34
  # Fallback provider keys (used only if the 3 mandatory vars above are incomplete)
35
  DEFAULT_BLAXEL_WORKSPACE = "vasanthfeb13"
@@ -118,7 +128,7 @@ def _normalize_token(value: str) -> str:
118
  return token
119
 
120
 
121
- def _build_client(api_base_url: str, hf_token: str) -> Any:
122
  if OpenAI is None:
123
  raise RuntimeError("openai package is not installed.")
124
 
@@ -128,8 +138,8 @@ def _build_client(api_base_url: str, hf_token: str) -> Any:
128
  default_headers["X-Blaxel-Workspace"] = workspace
129
 
130
  if default_headers:
131
- return OpenAI(api_key=_normalize_token(hf_token), base_url=api_base_url, default_headers=default_headers)
132
- return OpenAI(api_key=_normalize_token(hf_token), base_url=api_base_url)
133
 
134
  # ---------------------------------------------------------------------------
135
  # Provider / runtime config resolution
@@ -152,16 +162,21 @@ def _resolve_client() -> tuple[Any, str] | None:
152
  """Return (client, model_name) or None if nothing is configured."""
153
  api_base = (API_BASE_URL or "").strip()
154
  model = (MODEL_NAME or "").strip()
155
- token = (HF_TOKEN or "").strip()
156
 
157
- # Primary: all 3 mandatory vars set
158
  if api_base and model and token:
159
  try:
160
  return _build_client(api_base, token), model
161
  except Exception:
162
- pass
 
 
 
 
 
163
 
164
- # Fallback: Blaxel
165
  blaxel_key = os.getenv("BLAXEL_AUTHORIZATION", "").strip()
166
  if blaxel_key:
167
  m = model or os.getenv("BLAXEL_MODEL", DEFAULT_BLAXEL_MODEL).strip()
@@ -172,7 +187,7 @@ def _resolve_client() -> tuple[Any, str] | None:
172
  except Exception:
173
  pass
174
 
175
- # Fallback: Cerebras
176
  cerebras_key = os.getenv("CEREBRAS_API_KEY", "").strip()
177
  if cerebras_key:
178
  m = model or os.getenv("CEREBRAS_MODEL", DEFAULT_CEREBRAS_MODEL).strip()
 
3
 
4
  MANDATORY submission variables:
5
  - API_BASE_URL: OpenAI-compatible chat completions API base URL.
6
+ - API_KEY: token for the provided OpenAI-compatible proxy (preferred).
7
+ - HF_TOKEN: accepted compatibility alias for API_KEY.
8
  - MODEL_NAME: model identifier for inference.
 
9
 
10
  STDOUT FORMAT (mandatory):
11
  [START] task=<task_name> env=soc_triage_env model=<model_name>
 
27
  # ---------------------------------------------------------------------------
28
  # Mandatory env vars
29
  # ---------------------------------------------------------------------------
30
+ API_BASE_URL = os.getenv("API_BASE_URL", "").strip() or os.getenv("OPENAI_API_BASE_URL", "").strip()
31
+ MODEL_NAME = os.getenv("MODEL_NAME", "sandbox-openai").strip()
32
+ API_KEY = (
33
+ os.getenv("API_KEY", "").strip()
34
+ or os.getenv("HF_TOKEN", "").strip()
35
+ or os.getenv("OPENAI_API_KEY", "").strip()
36
+ )
37
  LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME")
38
+ ALLOW_PROVIDER_FALLBACK = os.getenv("ALLOW_PROVIDER_FALLBACK", "0").strip().lower() in {
39
+ "1",
40
+ "true",
41
+ "yes",
42
+ }
43
 
44
  # Fallback provider keys (used only if the 3 mandatory vars above are incomplete)
45
  DEFAULT_BLAXEL_WORKSPACE = "vasanthfeb13"
 
128
  return token
129
 
130
 
131
+ def _build_client(api_base_url: str, api_key: str) -> Any:
132
  if OpenAI is None:
133
  raise RuntimeError("openai package is not installed.")
134
 
 
138
  default_headers["X-Blaxel-Workspace"] = workspace
139
 
140
  if default_headers:
141
+ return OpenAI(api_key=_normalize_token(api_key), base_url=api_base_url, default_headers=default_headers)
142
+ return OpenAI(api_key=_normalize_token(api_key), base_url=api_base_url)
143
 
144
  # ---------------------------------------------------------------------------
145
  # Provider / runtime config resolution
 
162
  """Return (client, model_name) or None if nothing is configured."""
163
  api_base = (API_BASE_URL or "").strip()
164
  model = (MODEL_NAME or "").strip()
165
+ token = (API_KEY or "").strip()
166
 
167
+ # Primary: validator-injected proxy configuration (API_KEY or HF_TOKEN)
168
  if api_base and model and token:
169
  try:
170
  return _build_client(api_base, token), model
171
  except Exception:
172
+ return None
173
+
174
+ # In submission mode we intentionally do not use personal provider creds,
175
+ # because Phase 2 requires traffic through the provided proxy.
176
+ if not ALLOW_PROVIDER_FALLBACK:
177
+ return None
178
 
179
+ # Optional local-only fallback: Blaxel
180
  blaxel_key = os.getenv("BLAXEL_AUTHORIZATION", "").strip()
181
  if blaxel_key:
182
  m = model or os.getenv("BLAXEL_MODEL", DEFAULT_BLAXEL_MODEL).strip()
 
187
  except Exception:
188
  pass
189
 
190
+ # Optional local-only fallback: Cerebras
191
  cerebras_key = os.getenv("CEREBRAS_API_KEY", "").strip()
192
  if cerebras_key:
193
  m = model or os.getenv("CEREBRAS_MODEL", DEFAULT_CEREBRAS_MODEL).strip()
server/validator_tests.log CHANGED
@@ -31,3 +31,6 @@
31
  {"timestamp": 1775847814.208392, "method": "GET", "url": "http://testserver/tasks", "status": 200, "latency_sec": 0.0007, "request_body": ""}
32
  {"timestamp": 1775847814.2101681, "method": "POST", "url": "http://testserver/grader", "status": 200, "latency_sec": 0.0008, "request_body": "{\"task_id\":\"easy\",\"action\":{\"tool_name\":\"submit_verdict\",\"classification\":\"high\",\"recommended_action\":\"contain\",\"reasoning\":\"test\"}}"}
33
  {"timestamp": 1775847814.211788, "method": "GET", "url": "http://testserver/logs", "status": 200, "latency_sec": 0.0007, "request_body": ""}
 
 
 
 
31
  {"timestamp": 1775847814.208392, "method": "GET", "url": "http://testserver/tasks", "status": 200, "latency_sec": 0.0007, "request_body": ""}
32
  {"timestamp": 1775847814.2101681, "method": "POST", "url": "http://testserver/grader", "status": 200, "latency_sec": 0.0008, "request_body": "{\"task_id\":\"easy\",\"action\":{\"tool_name\":\"submit_verdict\",\"classification\":\"high\",\"recommended_action\":\"contain\",\"reasoning\":\"test\"}}"}
33
  {"timestamp": 1775847814.211788, "method": "GET", "url": "http://testserver/logs", "status": 200, "latency_sec": 0.0007, "request_body": ""}
34
+ {"timestamp": 1775850347.681291, "method": "POST", "url": "http://127.0.0.1:8002/reset", "status": 200, "latency_sec": 0.0023, "request_body": "{}"}
35
+ {"timestamp": 1775850409.608244, "method": "GET", "url": "http://127.0.0.1:8003/schema", "status": 200, "latency_sec": 0.0038, "request_body": ""}
36
+ {"timestamp": 1775850409.62133, "method": "GET", "url": "http://127.0.0.1:8003/tasks", "status": 200, "latency_sec": 0.001, "request_body": ""}