diff --git a/freerelay/core/routing/engine.py b/freerelay/core/routing/engine.py index 2a17456..8a4e909 100644 --- a/freerelay/core/routing/engine.py +++ b/freerelay/core/routing/engine.py @@ -92,12 +92,17 @@ class RequestContext: total_tokens: int schema_success_ratio: float tenant_tier: str + routing_preference: str = "balanced" policy_directive: RoutingDirective = field(default_factory=RoutingDirective) def policy_context(self) -> dict[str, Any]: return { "workload": self.workload_profile.to_dict(), - "tenant": {"id": self.user_id, "tier": self.tenant_tier}, + "tenant": { + "id": self.user_id, + "tier": self.tenant_tier, + "routing_preference": self.routing_preference + }, "schema": {"success_ratio": self.schema_success_ratio}, "compression": { "tokens_saved": self.compression.tokens_saved, @@ -206,6 +211,7 @@ def _prepare_context( request: ChatCompletionRequest, user_id: str | None = None, tier: str = "free", + routing_preference: str = "balanced", ) -> RequestContext: profile = self.profiler.profile(request) bundle = self.context_optimizer.optimize(request) @@ -221,6 +227,7 @@ def _prepare_context( total_tokens=bundle.total_tokens, schema_success_ratio=0.95, tenant_tier=tier, + routing_preference=routing_preference, ) def _build_policy_context(self, context: RequestContext) -> dict[str, Any]: @@ -334,8 +341,11 @@ async def route( request: ChatCompletionRequest, user_id: str | None = None, tier: str = "free", + routing_preference: str = "balanced", ) -> ChatCompletionResponse: - context = self._prepare_context(request, user_id=user_id, tier=tier) + context = self._prepare_context( + request, user_id=user_id, tier=tier, routing_preference=routing_preference + ) ranked, directive = await self._ranked_slots(context) if not ranked: @@ -439,8 +449,11 @@ async def route_stream( request: ChatCompletionRequest, user_id: str | None = None, tier: str = "free", + routing_preference: str = "balanced", ) -> Any: - context = self._prepare_context(request, user_id=user_id, tier=tier) + context = self._prepare_context( + request, user_id=user_id, tier=tier, routing_preference=routing_preference + ) ranked, _ = await self._ranked_slots(context) if not ranked: diff --git a/freerelay/core/routing/policy.py b/freerelay/core/routing/policy.py index dadb661..007668c 100644 --- a/freerelay/core/routing/policy.py +++ b/freerelay/core/routing/policy.py @@ -51,6 +51,7 @@ class RoutingDirective: require_hedging: str | None = None human_gate: bool = False policy_weight: float = 1.0 + routing_preference: str = "balanced" @dataclass @@ -84,6 +85,7 @@ def directive(self) -> RoutingDirective: require_hedging=self.require_hedging, human_gate=self.human_gate, policy_weight=self.policy_weight, + routing_preference="balanced", # Default, can be overridden by rule if we want ) @@ -149,19 +151,25 @@ def apply( available_providers: list[str], ) -> tuple[list[str], RoutingDirective]: """Reorder providers and surface the matching directive.""" + tenant_pref = context.get("tenant", {}).get("routing_preference", "balanced") + for rule in self.rules: if self._eval_condition(rule.condition, context): directive = rule.directive + directive.routing_preference = tenant_pref # Apply tenant preference logger.info( - "Routing rule %s matched → %s", + "Routing rule %s matched → %s (pref: %s)", rule.name, ", ".join(directive.prefer or ["(no prefer)"]), + tenant_pref, ) ordered = self._reorder_providers( available_providers, directive.prefer, directive.exclude ) return ordered, directive - return available_providers, RoutingDirective() + + default_directive = RoutingDirective(routing_preference=tenant_pref) + return available_providers, default_directive def _eval_condition(self, condition: str, context: dict[str, Any]) -> bool: return _eval_condition_context(condition, context) diff --git a/freerelay/core/routing/scorer.py b/freerelay/core/routing/scorer.py index eaf4a52..93d451d 100644 --- a/freerelay/core/routing/scorer.py +++ b/freerelay/core/routing/scorer.py @@ -113,7 +113,20 @@ def compute_expected_utility( safety = _safety_multiplier(profile, provider_models) policy_weight = directive.policy_weight if directive else 1.0 - return success_prob * quality * schema * latency * cost * safety * policy_weight + # Adjust weights based on routing preference + pref = directive.routing_preference if directive else "balanced" + if pref == "cost-optimized": + # Square cost to make it more dominant, root latency and quality + cost = cost**1.5 + latency = latency**0.5 + quality = quality**0.5 + elif pref == "performance-first": + # Square latency and quality, root cost + latency = latency**1.5 + quality = quality**1.5 + cost = cost**0.5 + + return success_prob * quality * latency * cost * safety * policy_weight * schema def compute_composite_score( diff --git a/freerelay/main.py b/freerelay/main.py index 0c34a9a..25c3362 100644 --- a/freerelay/main.py +++ b/freerelay/main.py @@ -38,6 +38,8 @@ CheckoutResponse, RegisterRequest, RegisterResponse, + TenantSettingsRequest, + TenantSettingsResponse, ) logger = logging.getLogger("freerelay") @@ -227,10 +229,11 @@ async def chat_completions(request: Request) -> Response: user_id = getattr(request.state, "user_id", None) tier = getattr(request.state, "tier", "free") + routing_preference = getattr(request.state, "routing_preference", "balanced") if req.is_streaming(): return StreamingResponse( - engine.route_stream(req, user_id=user_id, tier=tier), + engine.route_stream(req, user_id=user_id, tier=tier, routing_preference=routing_preference), media_type="text/event-stream", headers={ "Cache-Control": "no-cache", @@ -239,7 +242,7 @@ async def chat_completions(request: Request) -> Response: }, ) - response = await engine.route(req, user_id=user_id, tier=tier) + response = await engine.route(req, user_id=user_id, tier=tier, routing_preference=routing_preference) if "error" in response.model_dump(): return JSONResponse( @@ -302,6 +305,68 @@ async def register(req: RegisterRequest) -> RegisterResponse: content={"error": f"Registration failed: {str(e)}"}, ) # type: ignore + @app.post("/v1/tenant/settings", response_model=TenantSettingsResponse) + async def update_tenant_settings( + request: Request, + settings_req: TenantSettingsRequest + ) -> TenantSettingsResponse: + from freerelay.shared.tenancy.supabase import get_supabase_admin_client + + user_id = getattr(request.state, "user_id", None) + if not user_id or user_id == "admin": + return JSONResponse( + status_code=401, + content={"error": "Authentication required to update settings"} + ) # type: ignore + + if settings_req.routing_preference not in ["cost-optimized", "balanced", "performance-first"]: + return JSONResponse( + status_code=400, + content={"error": "Invalid routing preference. Must be 'cost-optimized', 'balanced', or 'performance-first'"} + ) # type: ignore + + try: + supabase = get_supabase_admin_client() + supabase.table("users").update( + {"routing_preference": settings_req.routing_preference} + ).eq("id", user_id).execute() + + return TenantSettingsResponse( + success=True, + routing_preference=settings_req.routing_preference + ) + except Exception as e: + logger.exception("Failed to update tenant settings") + return JSONResponse( + status_code=500, + content={"error": f"Failed to update settings: {str(e)}"} + ) # type: ignore + + @app.get("/v1/tenant/settings", response_model=TenantSettingsResponse) + async def get_tenant_settings(request: Request) -> TenantSettingsResponse: + from freerelay.shared.tenancy.supabase import get_supabase_client + + user_id = getattr(request.state, "user_id", None) + if not user_id or user_id == "admin": + return JSONResponse( + status_code=401, + content={"error": "Authentication required"} + ) # type: ignore + + try: + supabase = get_supabase_client() + result = supabase.table("users").select("routing_preference").eq("id", user_id).execute() + if result.data: + pref = result.data[0].get("routing_preference", "balanced") + return TenantSettingsResponse(success=True, routing_preference=pref) + return TenantSettingsResponse(success=False, routing_preference="balanced") + except Exception as e: + logger.exception("Failed to fetch tenant settings") + return JSONResponse( + status_code=500, + content={"error": f"Failed to fetch settings: {str(e)}"} + ) # type: ignore + @app.post("/v1/billing/checkout", response_model=None) async def billing_checkout(req: CheckoutRequest) -> CheckoutResponse: from freerelay.integrations.stripe import create_checkout_session diff --git a/freerelay/middleware/auth.py b/freerelay/middleware/auth.py index 1e028ca..053a010 100644 --- a/freerelay/middleware/auth.py +++ b/freerelay/middleware/auth.py @@ -33,7 +33,7 @@ def _verify_token_supabase(token_hash: str) -> dict[str, str] | None: # Join with users table to get the tier result = ( supabase.table("api_keys") - .select("user_id, users(tier)") + .select("user_id, users(tier, routing_preference)") .eq("key_hash", token_hash) .eq("is_active", True) .execute() @@ -43,9 +43,11 @@ def _verify_token_supabase(token_hash: str) -> dict[str, str] | None: user_id = str(data["user_id"]) users_data: Any = data.get("users") tier = "free" + routing_preference = "balanced" if isinstance(users_data, dict): tier = str(users_data.get("tier", "free")) - return {"user_id": user_id, "tier": tier} + routing_preference = str(users_data.get("routing_preference", "balanced")) + return {"user_id": user_id, "tier": tier, "routing_preference": routing_preference} return None except Exception as e: logger.error(f"Supabase auth error: {e}") @@ -86,6 +88,7 @@ async def dispatch( if self.api_key and token == self.api_key: request.state.user_id = "admin" request.state.tier = "gold" + request.state.routing_preference = "balanced" return await call_next(request) # 2. Supabase check @@ -95,6 +98,7 @@ async def dispatch( if user_info: request.state.user_id = user_info["user_id"] request.state.tier = user_info["tier"] + request.state.routing_preference = user_info["routing_preference"] return await call_next(request) return JSONResponse( diff --git a/freerelay/shared/models/internal.py b/freerelay/shared/models/internal.py index 3e22ed1..6fa022a 100644 --- a/freerelay/shared/models/internal.py +++ b/freerelay/shared/models/internal.py @@ -316,6 +316,13 @@ class RegisterRequest(BaseModel): class RegisterResponse(BaseModel): api_key: str +class TenantSettingsRequest(BaseModel): + routing_preference: str # 'cost-optimized', 'balanced', 'performance-first' + +class TenantSettingsResponse(BaseModel): + success: bool + routing_preference: str + class CheckoutRequest(BaseModel): email: str diff --git a/supabase_schema.sql b/supabase_schema.sql index aace945..7080af0 100644 --- a/supabase_schema.sql +++ b/supabase_schema.sql @@ -5,6 +5,7 @@ CREATE TABLE IF NOT EXISTS users ( id UUID PRIMARY KEY DEFAULT uuid_generate_v4(), email TEXT UNIQUE NOT NULL, tier TEXT NOT NULL DEFAULT 'free', -- 'free', 'bronze', 'silver', 'gold' + routing_preference TEXT NOT NULL DEFAULT 'balanced', -- 'cost-optimized', 'balanced', 'performance-first' created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW() );