# 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