Spaces:
Sleeping
Sleeping
| # 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, | |
| ) | |
| def state(self) -> State: | |
| return self._state |