architecture-env / server /architecture_env_environment.py
thepikachu's picture
Fixed architecture step function
abc93f5 verified
Raw
History Blame Contribute Delete
52.5 kB
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.
"""
Architecture Env Environment Implementation.
Problem-driven architecture synthesis environment.
The agent receives a software description and must incrementally:
- add components
- choose the right technology for each architectural concern
- connect components
- submit the final design
Scoring is deterministic and returns a reward in [0, 1].
Key design principles:
- Every component maps to an abstract family; the grader scores families, not exact names.
- "connect" parsing is fixed to handle multi-word component names reliably.
- Family fit score rewards satisfying a family with a good choice ONCE per family, capped.
- Any real-world technology can be added; unknown techs are accepted as generic "custom" items
(they don't satisfy family requirements but don't raise invalid_action either).
"""
from __future__ import annotations
import re
from uuid import uuid4
from typing import Any, Dict, List, Optional, Set, Tuple
from openenv.core.env_server.interfaces import Environment
from openenv.core.env_server.types import State
try:
from ..models import ArchitectureAction, ArchitectureObservation
except ImportError:
from models import ArchitectureAction, ArchitectureObservation
class ArchitectureEnvironment(Environment):
SUPPORTS_CONCURRENT_SESSIONS: bool = True
MAX_STEPS: int = 30 # Raised because richer tasks need more actions.
TASKS: List[Dict[str, Any]] = [
# ── EASY ──────────────────────────────────────────────────────────────
{
"name": "url_shortener",
"difficulty": "easy",
"prompt": (
"Design a scalable URL shortener backend. "
"Think about request routing, persistence, caching, abuse prevention, and observability."
),
"required_items": ["load_balancer", "api_server", "database", "cache"],
"bonus_items": ["rate_limiting", "observability", "auth"],
"required_connections": [
("load_balancer", "api_server"),
("api_server", "database"),
("api_server", "cache"),
],
"bonus_connections": [
("api_server", "rate_limiting"),
("api_server", "observability"),
],
"family_preferences": {
# Preferred techs per family for this task.
# Score is awarded per-family (not per-tech), so stacking techs
# in the same family does NOT inflate the score.
"database": {"preferred": ["postgres", "mysql", "cockroachdb"], "discouraged": ["sqlite"]},
"cache": {"preferred": ["redis", "memcached", "keydb"], "discouraged": []},
"load_balancer": {"preferred": ["nginx", "haproxy", "envoy", "traefik"], "discouraged": []},
"auth": {"preferred": ["oauth2", "keycloak", "auth0"], "discouraged": []},
"observability": {"preferred": ["prometheus", "grafana", "opentelemetry", "datadog"], "discouraged": []},
},
},
# ── MEDIUM ────────────────────────────────────────────────────────────
{
"name": "chat_system",
"difficulty": "medium",
"prompt": (
"Design a real-time chat system. "
"Think about websocket delivery, message fanout, persistence, presence, push notifications, and async processing."
),
"required_items": ["api_server", "websocket_gateway", "database", "cache", "broker", "worker"],
"bonus_items": ["presence_service", "notification_service", "observability", "auth"],
"required_connections": [
("api_server", "websocket_gateway"),
("websocket_gateway", "broker"),
("broker", "worker"),
("worker", "database"),
("worker", "cache"),
],
"bonus_connections": [
("worker", "presence_service"),
("worker", "notification_service"),
("api_server", "observability"),
("api_server", "auth"),
],
"family_preferences": {
# RabbitMQ/NATS are preferred for fanout; Kafka is overkill here.
"broker": {"preferred": ["rabbitmq", "nats"], "discouraged": ["activemq", "kafka"]},
"database": {"preferred": ["postgres", "mongodb", "cassandra"],"discouraged": ["sqlite"]},
"cache": {"preferred": ["redis"], "discouraged": []},
"auth": {"preferred": ["keycloak", "oauth2", "auth0"], "discouraged": []},
"observability": {"preferred": ["prometheus", "grafana", "opentelemetry"], "discouraged": []},
},
},
{
"name": "ecommerce_platform",
"difficulty": "medium",
"prompt": (
"Design an e-commerce platform backend. "
"Think about product catalog, user accounts, cart, orders, payments, inventory, and notifications."
),
"required_items": ["api_server", "database", "cache", "broker", "search"],
"bonus_items": ["payment_gateway", "notification_service", "observability", "auth", "rate_limiting", "cdn"],
"required_connections": [
("api_server", "database"),
("api_server", "cache"),
("api_server", "search"),
("api_server", "broker"),
("broker", "worker"),
],
"bonus_connections": [
("api_server", "payment_gateway"),
("api_server", "auth"),
("api_server", "observability"),
("worker", "notification_service"),
],
"family_preferences": {
"database": {"preferred": ["postgres", "mysql", "cockroachdb"], "discouraged": ["sqlite"]},
"cache": {"preferred": ["redis", "memcached"], "discouraged": []},
"broker": {"preferred": ["rabbitmq", "kafka", "pulsar"], "discouraged": ["activemq"]},
"search": {"preferred": ["elasticsearch", "opensearch"], "discouraged": []},
"auth": {"preferred": ["keycloak", "auth0", "oauth2"], "discouraged": []},
"observability": {"preferred": ["prometheus", "grafana", "datadog"], "discouraged": []},
},
},
# ── HARD ──────────────────────────────────────────────────────────────
{
"name": "youtube_platform",
"difficulty": "hard",
"prompt": (
"Design a YouTube-like video platform. "
"Think about uploads, object storage, transcoding pipelines, event streaming, search, "
"recommendations, CDN delivery, caching, and scalability."
),
"required_items": [
"api_server", "storage", "database", "cache",
"broker", "transcoder_worker", "cdn", "search",
],
"bonus_items": ["recommendation_worker", "observability", "rate_limiting", "auth", "metadata_service"],
"required_connections": [
("api_server", "storage"),
("api_server", "broker"),
("broker", "transcoder_worker"),
("transcoder_worker", "storage"),
("api_server", "cdn"),
("api_server", "search"),
],
"bonus_connections": [
("transcoder_worker", "database"),
("recommendation_worker", "cache"),
("recommendation_worker", "search"),
("api_server", "observability"),
("api_server", "rate_limiting"),
("api_server", "auth"),
],
"family_preferences": {
# Kafka/Pulsar preferred for high-throughput streaming; RabbitMQ is acceptable but not ideal.
"broker": {"preferred": ["kafka", "pulsar"], "discouraged": ["activemq", "rabbitmq"]},
"storage": {"preferred": ["s3", "gcs", "azure_blob", "minio", "object_storage"], "discouraged": []},
"database": {"preferred": ["postgres", "cockroachdb", "cassandra"], "discouraged": ["sqlite"]},
"cache": {"preferred": ["redis"], "discouraged": []},
"search": {"preferred": ["elasticsearch", "opensearch"], "discouraged": []},
"cdn": {"preferred": ["cloudfront", "cloudflare", "fastly", "akamai"], "discouraged": []},
"auth": {"preferred": ["keycloak", "auth0", "oauth2"], "discouraged": []},
"observability": {"preferred": ["prometheus", "grafana", "opentelemetry", "elk"], "discouraged": []},
},
},
{
"name": "ride_sharing",
"difficulty": "hard",
"prompt": (
"Design a ride-sharing platform like Uber. "
"Think about real-time location tracking, driver matching, trip state machines, "
"payments, surge pricing, notifications, and geo-spatial querying."
),
"required_items": [
"api_server", "websocket_gateway", "database", "cache",
"broker", "worker", "geospatial_index",
],
"bonus_items": ["payment_gateway", "notification_service", "observability", "auth", "rate_limiting", "search"],
"required_connections": [
("api_server", "websocket_gateway"),
("api_server", "geospatial_index"),
("api_server", "database"),
("api_server", "cache"),
("api_server", "broker"),
("broker", "worker"),
],
"bonus_connections": [
("api_server", "payment_gateway"),
("worker", "notification_service"),
("api_server", "auth"),
("api_server", "observability"),
],
"family_preferences": {
"database": {"preferred": ["postgres", "cockroachdb", "cassandra"], "discouraged": ["sqlite"]},
"cache": {"preferred": ["redis"], "discouraged": []},
"broker": {"preferred": ["kafka", "pulsar", "rabbitmq"], "discouraged": ["activemq"]},
"geospatial_index": {"preferred": ["postgis", "redis_geo", "elasticsearch"], "discouraged": []},
"auth": {"preferred": ["keycloak", "auth0", "oauth2"], "discouraged": []},
"observability": {"preferred": ["prometheus", "grafana", "datadog"], "discouraged": []},
},
},
{
"name": "ml_platform",
"difficulty": "hard",
"prompt": (
"Design an ML training and serving platform. "
"Think about data ingestion, feature stores, experiment tracking, model training, "
"model registry, online inference, batch inference, and observability."
),
"required_items": [
"api_server", "storage", "database", "broker",
"feature_store", "model_registry", "inference_server",
],
"bonus_items": ["experiment_tracker", "batch_worker", "observability", "auth", "cache", "search"],
"required_connections": [
("api_server", "storage"),
("api_server", "database"),
("api_server", "broker"),
("broker", "feature_store"),
("feature_store", "inference_server"),
("model_registry", "inference_server"),
],
"bonus_connections": [
("api_server", "experiment_tracker"),
("api_server", "auth"),
("api_server", "observability"),
("batch_worker", "storage"),
("batch_worker", "feature_store"),
],
"family_preferences": {
"storage": {"preferred": ["s3", "gcs", "azure_blob", "minio"], "discouraged": []},
"database": {"preferred": ["postgres", "mysql"], "discouraged": ["sqlite"]},
"broker": {"preferred": ["kafka", "pulsar"], "discouraged": ["activemq"]},
"feature_store": {"preferred": ["feast", "hopsworks", "tecton"], "discouraged": []},
"model_registry": {"preferred": ["mlflow", "neptune", "wandb"], "discouraged": []},
"inference_server": {"preferred": ["triton", "torchserve", "ray_serve", "bento_ml"], "discouraged": []},
"observability": {"preferred": ["prometheus", "grafana", "opentelemetry"], "discouraged": []},
},
},
]
# ── TASK FOCUS FAMILIES ──────────────────────────────────────────────────
# These are the families that are coherent for each task. Anything outside
# this set is treated as architectural drift and receives a deterministic
# penalty even if the tech is otherwise valid.
TASK_FAMILY_FOCUS: Dict[str, Set[str]] = {
"url_shortener": {
"load_balancer", "database", "cache", "auth", "observability", "rate_limiting",
},
"chat_system": {
"load_balancer", "websocket_gateway", "database", "cache", "broker", "worker",
"presence_service", "notification_service", "auth", "observability", "rate_limiting",
},
"ecommerce_platform": {
"load_balancer", "database", "cache", "broker", "search",
"payment_gateway", "notification_service", "auth", "observability",
"rate_limiting", "cdn", "worker",
},
"youtube_platform": {
"load_balancer", "storage", "database", "cache", "broker", "transcoder_worker",
"cdn", "search", "recommendation_worker", "metadata_service", "auth",
"observability", "rate_limiting",
},
"ride_sharing": {
"load_balancer", "websocket_gateway", "database", "cache", "broker", "worker",
"geospatial_index", "payment_gateway", "notification_service", "auth",
"observability", "rate_limiting", "search",
},
"ml_platform": {
"load_balancer", "storage", "database", "broker", "feature_store", "model_registry",
"inference_server", "experiment_tracker", "batch_worker", "auth",
"observability", "cache", "search",
},
}
# ── COMPONENT ALIASES ────────────────────────────────────────────────────
# Maps every raw user string β†’ canonical internal name.
# Canonical names that are also family member names resolve via FAMILY_GROUPS.
COMPONENT_ALIASES: Dict[str, str] = {
# ── API / Gateway
"api": "api_server",
"api server": "api_server",
"api_server": "api_server",
"backend": "api_server",
"rest api": "api_server",
"graphql api": "api_server",
"grpc server": "api_server",
"app server": "api_server",
"express": "api_server",
"fastapi": "api_server",
"django": "api_server",
"flask": "api_server",
"spring boot": "api_server",
"rails": "api_server",
"laravel": "api_server",
# ── Load Balancer
"load balancer": "load_balancer",
"load_balancer": "load_balancer",
"lb": "load_balancer",
"nginx": "nginx",
"haproxy": "haproxy",
"envoy": "envoy",
"traefik": "traefik",
"caddy": "caddy",
"aws alb": "load_balancer",
"gcp load balancer": "load_balancer",
# ── WebSocket / Realtime
"websocket": "websocket_gateway",
"websocket gateway": "websocket_gateway",
"websocket_gateway": "websocket_gateway",
"gateway": "websocket_gateway",
"socket io": "websocket_gateway",
"socket.io": "websocket_gateway",
"pusher": "websocket_gateway",
"ably": "websocket_gateway",
"centrifugo": "websocket_gateway",
"ws": "websocket_gateway",
# ── Databases (relational)
"database": "database",
"db": "database",
"rdbms": "database",
"postgres": "postgres",
"postgresql": "postgres",
"mysql": "mysql",
"mariadb": "mariadb",
"cockroachdb": "cockroachdb",
"tidb": "tidb",
"planetscale": "planetscale",
"aurora": "aurora",
"sqlite": "sqlite",
"neon": "neon",
"supabase": "supabase",
# ── Databases (document)
"mongodb": "mongodb",
"mongo": "mongodb",
"firestore": "firestore",
"couchdb": "couchdb",
"dynamodb": "dynamodb",
# ── Databases (wide column)
"cassandra": "cassandra",
"scylladb": "scylladb",
"hbase": "hbase",
"bigtable": "bigtable",
# ── Databases (time-series)
"timescaledb": "timescaledb",
"influxdb": "influxdb",
"prometheus tsdb": "timescaledb",
"victoriametrics": "victoriametrics",
"questdb": "questdb",
# ── Databases (graph)
"neo4j": "neo4j",
"janusgraph": "janusgraph",
"neptune": "neptune",
"tigergraph": "tigergraph",
"dgraph": "dgraph",
# ── Databases (vector)
"pinecone": "pinecone",
"weaviate": "weaviate",
"milvus": "milvus",
"qdrant": "qdrant",
"chroma": "chroma",
"pgvector": "pgvector",
# ── Cache
"cache": "cache",
"redis": "redis",
"memcached": "memcached",
"keydb": "keydb",
"hazelcast": "hazelcast",
"dragonfly": "dragonfly",
"momento": "momento",
# ── Message Brokers / Queues
"broker": "broker",
"queue": "queue",
"message queue": "broker",
"message broker": "broker",
"rabbitmq": "rabbitmq",
"kafka": "kafka",
"pulsar": "pulsar",
"activemq": "activemq",
"nats": "nats",
"sqs": "sqs",
"pubsub": "pubsub",
"google pubsub": "pubsub",
"azure service bus": "azure_service_bus",
"azure_service_bus": "azure_service_bus",
"sns": "sns",
"kinesis": "kinesis",
"eventbridge": "eventbridge",
"redpanda": "redpanda",
# ── Stream Processing
"flink": "flink",
"spark": "spark",
"spark streaming": "spark",
"storm": "storm",
"samza": "samza",
"beam": "beam",
# ── Workers / Async
"worker": "worker",
"background worker": "worker",
"celery": "celery",
"sidekiq": "sidekiq",
"temporal": "temporal",
"cadence": "cadence",
"airflow": "airflow",
"prefect": "prefect",
"dagster": "dagster",
"transcoder": "transcoder_worker",
"transcoder worker": "transcoder_worker",
"transcoder_worker": "transcoder_worker",
"ffmpeg worker": "transcoder_worker",
"recommendation": "recommendation_worker",
"recommendation worker": "recommendation_worker",
"recommendation_worker": "recommendation_worker",
"batch worker": "batch_worker",
"batch_worker": "batch_worker",
# ── Object Storage
"storage": "storage",
"object storage": "object_storage",
"object_storage": "object_storage",
"s3": "s3",
"gcs": "gcs",
"google cloud storage": "gcs",
"azure blob": "azure_blob",
"azure_blob": "azure_blob",
"minio": "minio",
"hdfs": "hdfs",
"ceph": "ceph",
"backblaze": "backblaze",
# ── CDN
"cdn": "cdn",
"cloudfront": "cloudfront",
"cloudflare": "cloudflare",
"fastly": "fastly",
"akamai": "akamai",
"bunny cdn": "bunny_cdn",
"bunny_cdn": "bunny_cdn",
# ── Search
"search": "search",
"search index": "search_index",
"search_index": "search_index",
"elasticsearch": "elasticsearch",
"opensearch": "opensearch",
"solr": "solr",
"meilisearch": "meilisearch",
"typesense": "typesense",
"algolia": "algolia",
"sphinx": "sphinx",
# ── Auth / Identity
"auth": "auth",
"authentication": "auth",
"authorization": "auth",
"identity": "auth",
"oauth2": "oauth2",
"keycloak": "keycloak",
"auth0": "auth0",
"cognito": "cognito",
"okta": "okta",
"clerk": "clerk",
"supertokens": "supertokens",
"passport": "passport",
"zitadel": "zitadel",
"ory": "ory",
# ── Observability
"monitoring": "monitoring",
"observability": "observability",
"prometheus": "prometheus",
"grafana": "grafana",
"opentelemetry": "opentelemetry",
"otel": "opentelemetry",
"elk": "elk",
"datadog": "datadog",
"newrelic": "newrelic",
"new relic": "newrelic",
"jaeger": "jaeger",
"zipkin": "zipkin",
"loki": "loki",
"splunk": "splunk",
"sentry": "sentry",
"dynatrace": "dynatrace",
"instana": "instana",
"cloudwatch": "cloudwatch",
"signalflow": "signalflow",
# ── Rate Limiting / API Gateway
"rate limiting": "rate_limiting",
"rate_limiting": "rate_limiting",
"api gateway": "api_gateway",
"api_gateway": "api_gateway",
"kong": "kong",
"apigee": "apigee",
"aws api gateway": "api_gateway",
"traefik gateway": "api_gateway",
# ── Presence / Notification
"presence": "presence_service",
"presence service": "presence_service",
"presence_service": "presence_service",
"notification": "notification_service",
"notification service": "notification_service",
"notification_service": "notification_service",
"fcm": "notification_service",
"apns": "notification_service",
"twilio": "notification_service",
"sendgrid": "notification_service",
"ses": "notification_service",
# ── Metadata
"metadata": "metadata_service",
"metadata service": "metadata_service",
"metadata_service": "metadata_service",
# ── Payment
"payment": "payment_gateway",
"payment gateway": "payment_gateway",
"payment_gateway": "payment_gateway",
"stripe": "payment_gateway",
"braintree": "payment_gateway",
"paypal": "payment_gateway",
"adyen": "payment_gateway",
# ── Geospatial
"geospatial": "geospatial_index",
"geospatial index": "geospatial_index",
"geospatial_index": "geospatial_index",
"postgis": "postgis",
"redis geo": "redis_geo",
"redis_geo": "redis_geo",
"h3": "h3",
# ── Service Mesh / Config
"service mesh": "service_mesh",
"service_mesh": "service_mesh",
"istio": "istio",
"linkerd": "linkerd",
"consul": "consul",
"etcd": "etcd",
"vault": "vault",
"zookeeper": "zookeeper",
# ── ML / Feature / Model components
"feature store": "feature_store",
"feature_store": "feature_store",
"feast": "feast",
"hopsworks": "hopsworks",
"tecton": "tecton",
"model registry": "model_registry",
"model_registry": "model_registry",
"mlflow": "mlflow",
"neptune": "neptune",
"wandb": "wandb",
"inference server": "inference_server",
"inference_server": "inference_server",
"triton": "triton",
"torchserve": "torchserve",
"ray serve": "ray_serve",
"ray_serve": "ray_serve",
"bento ml": "bento_ml",
"bento_ml": "bento_ml",
"experiment tracker": "experiment_tracker",
"experiment_tracker": "experiment_tracker",
"mlflow tracking": "experiment_tracker",
}
# ── FAMILY GROUPS ────────────────────────────────────────────────────────
# Maps abstract family name β†’ all concrete tech members.
# Requirements use family names; a requirement is satisfied when ANY member is present.
FAMILY_GROUPS: Dict[str, Set[str]] = {
"load_balancer": {
"load_balancer", "nginx", "haproxy", "envoy", "traefik", "caddy",
},
"database": {
"database",
# Relational
"postgres", "mysql", "mariadb", "cockroachdb", "tidb", "planetscale",
"aurora", "sqlite", "neon", "supabase",
# Document
"mongodb", "firestore", "couchdb", "dynamodb",
# Wide column
"cassandra", "scylladb", "hbase", "bigtable",
# Time series
"timescaledb", "influxdb", "victoriametrics", "questdb",
# Graph
"neo4j", "janusgraph", "neptune", "tigergraph", "dgraph",
},
"cache": {
"cache", "redis", "memcached", "keydb", "hazelcast", "dragonfly", "momento",
},
"broker": {
"broker", "queue", "rabbitmq", "kafka", "pulsar", "activemq", "nats",
"sqs", "pubsub", "azure_service_bus", "sns", "kinesis", "eventbridge",
"redpanda",
},
"storage": {
"storage", "object_storage", "s3", "gcs", "azure_blob", "minio",
"hdfs", "ceph", "backblaze",
},
"cdn": {
"cdn", "cloudfront", "cloudflare", "fastly", "akamai", "bunny_cdn",
},
"search": {
"search", "search_index", "elasticsearch", "opensearch", "solr",
"meilisearch", "typesense", "algolia", "sphinx",
},
"auth": {
"auth", "oauth2", "keycloak", "auth0", "cognito", "okta", "clerk",
"supertokens", "passport", "zitadel", "ory",
},
"observability": {
"monitoring", "observability", "prometheus", "grafana", "opentelemetry",
"elk", "datadog", "newrelic", "jaeger", "zipkin", "loki", "splunk",
"sentry", "dynatrace", "instana", "cloudwatch", "signalflow",
},
"rate_limiting": {"rate_limiting", "kong", "apigee", "api_gateway"},
"notification_service": {
"notification_service", "fcm", "apns", "twilio", "sendgrid", "ses",
},
"payment_gateway": {
"payment_gateway", "stripe", "braintree", "paypal", "adyen",
},
"geospatial_index": {
"geospatial_index", "postgis", "redis_geo", "h3",
},
"feature_store": {"feature_store", "feast", "hopsworks", "tecton"},
"model_registry": {"model_registry", "mlflow", "neptune", "wandb"},
"inference_server": {
"inference_server", "triton", "torchserve", "ray_serve", "bento_ml",
},
"experiment_tracker": {"experiment_tracker"},
"websocket_gateway": {
"websocket_gateway", "socket_io", "pusher", "ably", "centrifugo",
},
# Stream processors satisfy a "stream_processor" family (for future tasks)
"stream_processor": {"flink", "spark", "storm", "samza", "beam"},
}
def __init__(self):
self._state = State(episode_id=str(uuid4()), step_count=0)
self._reset_count = 0
self._task_index = -1
self._arch_state: Dict[str, Any] = {}
self._start_new_task()
# ── Internal helpers ──────────────────────────────────────────────────────
def _start_new_task(self) -> None:
self._task_index = (self._task_index + 1) % len(self.TASKS)
task = self.TASKS[self._task_index]
self._arch_state = {
"episode_id": str(uuid4()),
"step_count": 0,
"task_name": task["name"],
"difficulty": task["difficulty"],
"prompt": task["prompt"],
"required_items": list(task["required_items"]),
"bonus_items": list(task["bonus_items"]),
"required_connections": [tuple(x) for x in task["required_connections"]],
"bonus_connections": [tuple(x) for x in task["bonus_connections"]],
"family_preferences": task.get("family_preferences", {}),
"components": set(),
"connections": set(),
"history": [],
"submitted": False,
"invalid_actions": 0,
"done": False,
"last_score": 0.0,
}
self._state = State(episode_id=self._arch_state["episode_id"], step_count=0)
def _normalize_text(self, text: str) -> str:
return re.sub(r"\s+", " ", (text or "").strip().lower())
def _canonicalize(self, raw: str) -> Optional[str]:
text = self._normalize_text(raw).replace("-", " ")
if text in self.COMPONENT_ALIASES:
return self.COMPONENT_ALIASES[text]
compact = text.replace(" ", "_")
if compact in self.COMPONENT_ALIASES:
return self.COMPONENT_ALIASES[compact]
# Partial match: if the token is a substring of an alias (e.g. "postgres db" β†’ "postgres")
for alias, canonical in self.COMPONENT_ALIASES.items():
if alias in text or text in alias:
return canonical
return None
def _family_of(self, component: str) -> Optional[str]:
for family, members in self.FAMILY_GROUPS.items():
if component in members:
return family
return None
def _options_for_requirement(self, req: str) -> Set[str]:
if req in self.FAMILY_GROUPS:
return set(self.FAMILY_GROUPS[req])
return {req}
def _requirement_satisfied(self, req: str, components: Set[str]) -> bool:
return bool(self._options_for_requirement(req) & components)
def _edge_satisfied(
self, edge: Tuple[str, str], components: Set[str], connections: Set[Tuple[str, str]]
) -> bool:
src_req, dst_req = edge
src_options = self._options_for_requirement(src_req)
dst_options = self._options_for_requirement(dst_req)
for s in src_options & components:
for d in dst_options & components:
if (s, d) in connections:
return True
return False
def _task_focus_families(self, task_name: Optional[str] = None) -> Set[str]:
task = (task_name or self._arch_state.get("task_name") or "").strip().lower()
return set(self.TASK_FAMILY_FOCUS.get(task, set()))
# ── FIX: parse "connect X Y" reliably for multi-word component names ─────
def _parse_message(self, message: str) -> Tuple[str, Any]:
text = self._normalize_text(message)
if not text:
return "noop", None
if text in {"submit", "done", "finish", "finalize", "finalise"}:
return "submit", None
add_match = re.match(r"^(?:add|include|put|place)\s+(.+)$", text)
if add_match:
return "add", add_match.group(1).strip()
remove_match = re.match(r"^(?:remove|delete|drop)\s+(.+)$", text)
if remove_match:
return "remove", remove_match.group(1).strip()
# FIX: For connect, try all split points to find a valid pair.
# This handles multi-word names like "api server websocket gateway" correctly.
connect_match = re.match(r"^connect\s+(.+)$", text)
if connect_match:
rest = connect_match.group(1).strip()
# Remove optional "to" connector word: "connect X to Y" β†’ "X Y"
rest = re.sub(r"\s+to\s+", " ", rest)
result = self._best_connect_split(rest)
if result:
return "connect", result
return "connect", (rest, "") # Will fail gracefully in _handle_connect
return "unknown", text
def _best_connect_split(self, rest: str) -> Optional[Tuple[str, str]]:
"""
Split 'rest' into two canonicalizable parts by trying every word boundary.
Returns the first split where BOTH parts are valid canonical names.
Strategy: try ALL splits and score each by (left_words + right_words) specificity,
preferring the split where both sides have the longest combined alias match.
This correctly resolves cases like "worker presence service" β†’ ("worker", "presence_service")
even though "worker presence" is also a partial alias substring.
Concretely: we try shortest-left-side first (i=1,2,...) so single-word left names
like "worker" are found before two-word false-positives like "worker presence".
"""
words = rest.split()
best: Optional[Tuple[str, str]] = None
best_score = -1
for i in range(1, len(words)):
left_raw = " ".join(words[:i])
right_raw = " ".join(words[i:])
left = self._canonicalize(left_raw)
right = self._canonicalize(right_raw)
if left and right:
# Score: prefer exact alias matches over partial/substring matches.
# An exact alias match gives len(words) bonus; partial gives 0.
left_exact = (
left_raw in self.COMPONENT_ALIASES
or left_raw.replace(" ", "_") in self.COMPONENT_ALIASES
)
right_exact = (
right_raw in self.COMPONENT_ALIASES
or right_raw.replace(" ", "_") in self.COMPONENT_ALIASES
)
split_score = int(left_exact) + int(right_exact)
if split_score > best_score:
best_score = split_score
best = (left, right)
return best
def _handle_add(self, payload: str) -> str:
component = self._canonicalize(payload)
if component is None:
# Accept unknown techs as custom components rather than penalizing.
# They won't satisfy family requirements but won't waste the agent's budget.
component = self._normalize_text(payload).replace(" ", "_")
if not component:
self._arch_state["invalid_actions"] += 1
return "Empty component name."
# Mark as custom (no family) by using the raw slug.
before = len(self._arch_state["components"])
self._arch_state["components"].add(component)
after = len(self._arch_state["components"])
if after == before:
self._arch_state["invalid_actions"] += 1
return f"Component '{component}' already present."
return f"Added component '{component}'."
def _handle_remove(self, payload: str) -> str:
component = self._canonicalize(payload)
if component is None:
component = self._normalize_text(payload).replace(" ", "_")
if component in self._arch_state["components"]:
self._arch_state["components"].remove(component)
self._arch_state["connections"] = {
edge for edge in self._arch_state["connections"] if component not in edge
}
return f"Removed component '{component}'."
self._arch_state["invalid_actions"] += 1
return f"Component '{component}' was not present."
def _handle_connect(self, payload: Tuple[str, str]) -> str:
left_raw, right_raw = payload
left = self._canonicalize(left_raw) or left_raw
right = self._canonicalize(right_raw) or right_raw
if not left or not right:
self._arch_state["invalid_actions"] += 1
return "Invalid connection: could not identify both components."
if left not in self._arch_state["components"] or right not in self._arch_state["components"]:
self._arch_state["invalid_actions"] += 1
return (
f"Both components must exist before connecting. "
f"Present: {sorted(self._arch_state['components'])}."
)
if left == right:
self._arch_state["invalid_actions"] += 1
return "Cannot connect a component to itself."
before = len(self._arch_state["connections"])
self._arch_state["connections"].add((left, right))
after = len(self._arch_state["connections"])
if after == before:
self._arch_state["invalid_actions"] += 1
return f"Connection '{left} -> {right}' already present."
return f"Connected {left} -> {right}."
# ── DETERMINISTIC SCORING ─────────────────────────────────────────────────
def _family_fit_score(self, components: Set[str]) -> Tuple[float, List[str], List[str]]:
"""
For each family that has preferences defined for this task, award:
+0.08 if the agent used a PREFERRED tech from that family
+0.04 if the agent used only ACCEPTABLE techs from that family
-0.05 for each DISCOURAGED tech used (penalty, not per-family-capped)
The preferred/acceptable bonus is capped per-family: using three preferred
techs from the same family still only gives +0.08, not +0.24.
This makes scoring independent of how many techs the agent stacks in one family.
"""
prefs = self._arch_state.get("family_preferences", {})
preferred_hits: List[str] = []
discouraged_hits: List[str] = []
score = 0.0
for family, conf in prefs.items():
preferred = set(conf.get("preferred", []))
acceptable = set(conf.get("acceptable", []))
discouraged = set(conf.get("discouraged", []))
family_members = self.FAMILY_GROUPS.get(family, set())
present_in_family = components & family_members
# Penalties are per-tech, not per-family (adding two discouraged techs is worse).
for tech in present_in_family:
if tech in discouraged:
score -= 0.05
discouraged_hits.append(tech)
# Bonus is per-family (capped): does the agent have at least one preferred/acceptable?
has_preferred = bool(present_in_family & preferred)
has_acceptable = bool(present_in_family & acceptable) and not has_preferred
if has_preferred:
score += 0.08
preferred_hits.extend(sorted(present_in_family & preferred))
elif has_acceptable:
score += 0.04
return score, preferred_hits, discouraged_hits
def _score_design(self) -> Tuple[float, Dict[str, Any]]:
required_items = set(self._arch_state["required_items"])
bonus_items = set(self._arch_state["bonus_items"])
required_connections = set(self._arch_state["required_connections"])
bonus_connections = set(self._arch_state["bonus_connections"])
components = set(self._arch_state["components"])
connections = set(self._arch_state["connections"])
task_name = self._arch_state["task_name"]
focus_families = self._task_focus_families(task_name)
matched_required_items = sorted(
[req for req in required_items if self._requirement_satisfied(req, components)]
)
missing_required_items = sorted([req for req in required_items if req not in matched_required_items])
matched_bonus_items = sorted(
[req for req in bonus_items if self._requirement_satisfied(req, components)]
)
matched_required_connections = sorted(
[edge for edge in required_connections if self._edge_satisfied(edge, components, connections)]
)
missing_required_connections = sorted(
[edge for edge in required_connections if edge not in matched_required_connections]
)
matched_bonus_connections = sorted(
[edge for edge in bonus_connections if self._edge_satisfied(edge, components, connections)]
)
component_coverage = len(matched_required_items) / max(1, len(required_items))
connection_coverage = len(matched_required_connections) / max(1, len(required_connections))
# Base score: components 42%, connections 30%.
score = 0.0
score += 0.42 * component_coverage
score += 0.30 * connection_coverage
# Bonus items and connections: capped totals.
score += min(0.12, 0.03 * len(matched_bonus_items))
score += min(0.08, 0.02 * len(matched_bonus_connections))
# Family-fit score: deterministic per-family bonuses/penalties.
# Cap the positive contribution to 0.16 so technology-fit bonus stays
# bounded regardless of how many family_preferences the task defines.
fit_bonus, preferred_hits, discouraged_hits = self._family_fit_score(components)
fit_bonus = max(-0.30, min(0.16, fit_bonus))
score += fit_bonus
# Coherence penalty: reward the architecture pattern, not just raw coverage.
# Any valid tech that lives outside the task's focused families should count
# as over-engineering or drift, even if it is a real technology.
irrelevant_components: List[str] = []
irrelevant_families: Set[str] = set()
coherence_penalty = 0.0
for component in components:
if component in required_items or component in bonus_items:
continue
family = self._family_of(component)
if family is None:
# Unknown custom tech is allowed, but it should not be "free".
# Keep the penalty small so custom stack choices remain possible.
if component.startswith("custom_"):
continue
coherence_penalty += 0.01
irrelevant_components.append(component)
continue
if family not in focus_families:
coherence_penalty += 0.035
irrelevant_components.append(component)
irrelevant_families.add(family)
# Extra slack for many unrelated families, even if each one is only used once.
if len(irrelevant_families) > 1:
coherence_penalty += 0.01 * (len(irrelevant_families) - 1)
# Extra slack for very large component sets that exceed the task pattern.
expected_baseline = len(required_items) + len(bonus_items)
if len(components) > expected_baseline + 2:
coherence_penalty += min(0.08, 0.004 * (len(components) - expected_baseline - 2))
score -= min(0.22, coherence_penalty)
# Penalty: components that are unrecognized (not in any family, not in bonus/required).
# Known-but-wrong-family components do NOT get penalized here; the discouraged penalty handles that.
all_known = set()
for members in self.FAMILY_GROUPS.values():
all_known |= members
all_known |= {c for items in [required_items, bonus_items] for c in items}
extra_items = [c for c in components if c not in all_known]
score -= min(0.10, 0.02 * len(extra_items))
# Invalid action penalty.
score -= min(0.08, 0.01 * self._arch_state["invalid_actions"])
score = max(0.0, min(1.0, score))
details = {
"task_name": self._arch_state["task_name"],
"difficulty": self._arch_state["difficulty"],
"prompt": self._arch_state["prompt"],
"required_items": sorted(required_items),
"bonus_items": sorted(bonus_items),
"present_components": sorted(components),
"matched_required_items": matched_required_items,
"missing_required_items": missing_required_items,
"matched_bonus_items": matched_bonus_items,
"required_connections": [list(x) for x in sorted(required_connections)],
"present_connections": [list(x) for x in sorted(connections)],
"matched_required_connections": [list(x) for x in matched_required_connections],
"missing_required_connections": [list(x) for x in missing_required_connections],
"matched_bonus_connections": [list(x) for x in matched_bonus_connections],
"preferred_hits": preferred_hits,
"discouraged_hits": discouraged_hits,
"irrelevant_components": sorted(irrelevant_components),
"irrelevant_families": sorted(irrelevant_families),
"coherence_penalty": round(coherence_penalty, 3),
"component_coverage": round(component_coverage, 3),
"connection_coverage": round(connection_coverage, 3),
"invalid_actions": self._arch_state["invalid_actions"],
"step_count": self._arch_state["step_count"],
"submitted": self._arch_state["submitted"],
"score": round(score, 3),
}
return score, details
def _build_observation(
self,
*,
action_message: str,
reward: float,
done: bool,
response_message: str,
details: Dict[str, Any],
) -> ArchitectureObservation:
present_components = details["present_components"]
present_connections = details["present_connections"]
conn_text = (
", ".join([f"{a}->{b}" for a, b in present_connections])
if present_connections
else "none"
)
summary = (
f"{self._arch_state['difficulty'].title()} task: {self._arch_state['task_name']}. "
f"{response_message} "
f"Components: {', '.join(present_components) if present_components else 'none'}. "
f"Connections: {conn_text}. "
f"Score: {details['score']:.3f}."
).strip()
metadata = dict(details)
metadata["last_action"] = action_message
metadata["history"] = list(self._arch_state["history"])
return ArchitectureObservation(
echoed_message=summary,
message_length=len(action_message or ""),
done=done,
reward=reward,
metadata=metadata,
)
# ── Public API ────────────────────────────────────────────────────────────
def reset(self) -> ArchitectureObservation:
self._reset_count += 1
self._start_new_task()
details = {
"task_name": self._arch_state["task_name"],
"difficulty": self._arch_state["difficulty"],
"prompt": self._arch_state["prompt"],
"required_items": list(self._arch_state["required_items"]),
"bonus_items": list(self._arch_state["bonus_items"]),
"present_components": [],
"matched_required_items": [],
"missing_required_items": list(self._arch_state["required_items"]),
"matched_bonus_items": [],
"required_connections": [list(x) for x in self._arch_state["required_connections"]],
"present_connections": [],
"matched_required_connections": [],
"missing_required_connections": [list(x) for x in self._arch_state["required_connections"]],
"matched_bonus_connections": [],
"preferred_hits": [],
"discouraged_hits": [],
"irrelevant_components": [],
"irrelevant_families": [],
"coherence_penalty": 0.0,
"component_coverage": 0.0,
"connection_coverage": 0.0,
"invalid_actions": 0,
"step_count": 0,
"submitted": False,
"score": 0.0,
"reset_count": self._reset_count,
}
return ArchitectureObservation(
echoed_message=(
f"Task: {self._arch_state['task_name']}. "
f"{self._arch_state['prompt']} "
"Use commands like 'add api server', 'add postgres', 'add rabbitmq', "
"'connect api server postgres', and 'submit'."
),
message_length=0,
done=False,
reward=0.0,
metadata=details,
)
def step(self, action: ArchitectureAction) -> ArchitectureObservation: # type: ignore[override]
if self._arch_state["done"]:
score, details = self._score_design()
return self._build_observation(
action_message=getattr(action, "message", ""),
reward=score,
done=True,
response_message="Episode already finished.",
details=details,
)
raw_message = getattr(action, "message", None) or getattr(action, "value", "") or ""
command, payload = self._parse_message(raw_message)
if command == "add":
response_message = self._handle_add(payload)
elif command == "remove":
response_message = self._handle_remove(payload)
elif command == "connect":
response_message = self._handle_connect(payload)
elif command == "submit":
self._arch_state["submitted"] = True
response_message = "Submission received."
elif command == "noop":
self._arch_state["invalid_actions"] += 1
response_message = "Empty action."
else:
self._arch_state["invalid_actions"] += 1
response_message = "Unrecognized command."
self._arch_state["step_count"] += 1
self._state.step_count = self._arch_state["step_count"]
score, details = self._score_design()
done = bool(
self._arch_state["submitted"]
or self._arch_state["step_count"] >= self.MAX_STEPS
)
self._arch_state["done"] = done
reward = score
if self._arch_state["submitted"] and score >= 0.80:
reward = min(1.0, score + 0.05)
self._arch_state["last_score"] = reward
self._arch_state["history"].append(
{
"step": self._arch_state["step_count"],
"action": raw_message,
"command": command,
"payload": payload if not isinstance(payload, tuple) else list(payload),
"score": round(reward, 3),
}
)
details["history"] = list(self._arch_state["history"])
details["submitted"] = self._arch_state["submitted"]
details["step_count"] = self._arch_state["step_count"]
details["score"] = round(reward, 3)
return self._build_observation(
action_message=raw_message,
reward=reward,
done=done,
response_message=response_message,
details=details,
)
@property
def state(self) -> State:
return self._state