Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
65 changes: 61 additions & 4 deletions backend/agents/create_agent_info.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
resolve_capacity,
)
from nexent.core.models.capacity_budget import (
BudgetResolverError,
RequestBudgetOverrides,
SafeInputBudgetCalculator,
UncertaintyReserveBasisUnknown,
Expand Down Expand Up @@ -80,7 +81,7 @@
NEXENT_SANDBOX_WORKSPACE_VOLUME,
)
from consts.model import ToolParamsRequest
from consts.exceptions import ValidationError
from consts.exceptions import ModelCapacityConfigError, ValidationError

logger = logging.getLogger("create_agent_info")
logger.setLevel(logging.INFO)
Expand Down Expand Up @@ -310,6 +311,28 @@ def _resolve_safe_input_budget(
exc,
)
return None
except BudgetResolverError as exc:
reason_by_type = {
"InvalidReservePolicy": "invalid_reserve_policy",
"RequestedOutputExceedsCapacity": "requested_output_exceeds_model",
"ReserveExceedsCapacity": "reserve_exceeds_capacity",
"NoSafeInputCapacity": "no_safe_input_capacity",
"SafeInputBudgetFingerprintMismatch": "budget_fingerprint_mismatch",
"CallerMaxTokensOverrideForbidden": "caller_output_override_forbidden",
"SafeInputBudgetCapacityMismatch": "capacity_snapshot_mismatch",
}
reason = reason_by_type.get(type(exc).__name__, "budget_resolution_failed")
logger.warning(
"W2 safe input budget rejected: tenant_id=%s model=%s reason=%s",
tenant_id,
capacity_snapshot.model_name,
reason,
)
raise ModelCapacityConfigError(
f"capacity_config_invalid.{reason}",
"The selected model capacity cannot produce a safe Agent input budget. "
"Review the model context, input, output, and reserve settings.",
) from exc
logger.debug(
"W2 safe input budget resolved: tenant_id=%s model=%s requested_output_tokens=%s "
"soft_input_budget_tokens=%s hard_input_budget_tokens=%s fingerprint=%s warnings=%s",
Expand Down Expand Up @@ -340,6 +363,12 @@ def _resolve_input_budget(
provider_raw = model_info.get("model_factory")
provider = provider_raw.lower().strip() if isinstance(provider_raw, str) else ""
model_id = model_info.get("model_name") or ""
persisted_profile_version = model_info.get("capability_profile_version")
if persisted_profile_version:
for (catalog_provider, catalog_model), profile in CAPABILITY_CATALOG.items():
if profile.capability_profile_version == persisted_profile_version:
provider, model_id = catalog_provider, catalog_model
break
provider_missing_detail = None
if not provider:
provider_missing_detail = (
Expand Down Expand Up @@ -880,7 +909,11 @@ async def create_model_config_list(tenant_id):
default_output_reserve_tokens=record.get("default_output_reserve_tokens"),
tokenizer_family=record.get("tokenizer_family"),
capacity_source=record.get("capacity_source"),
capability_profile_version=record.get("capability_profile_version")))
capability_profile_version=record.get("capability_profile_version"),
canonical_model_id=record.get("canonical_model_id"),
model_identity_metadata=record.get("model_identity_metadata"),
tokenizer_match_metadata=record.get("tokenizer_match_metadata"),
token_count_probe_metadata=record.get("token_count_probe_metadata")))
# fit for old version, main_model and sub_model use default model
main_model_config = tenant_config_manager.get_model_config(
key=MODEL_CONFIG_MAPPING["llm"], tenant_id=tenant_id)
Expand All @@ -896,7 +929,19 @@ async def create_model_config_list(tenant_id):
model_factory=main_model_config.get("model_factory"),
timeout_seconds=main_model_config.get("timeout_seconds"),
concurrency_limit=main_model_config.get("concurrency_limit"),
prompt_cache=main_prompt_cache))
prompt_cache=main_prompt_cache,
max_output_tokens=main_model_config.get("max_output_tokens"),
max_tokens=main_model_config.get("max_tokens"),
context_window_tokens=main_model_config.get("context_window_tokens"),
max_input_tokens=main_model_config.get("max_input_tokens"),
default_output_reserve_tokens=main_model_config.get("default_output_reserve_tokens"),
tokenizer_family=main_model_config.get("tokenizer_family"),
capacity_source=main_model_config.get("capacity_source"),
capability_profile_version=main_model_config.get("capability_profile_version"),
canonical_model_id=main_model_config.get("canonical_model_id"),
model_identity_metadata=main_model_config.get("model_identity_metadata"),
tokenizer_match_metadata=main_model_config.get("tokenizer_match_metadata"),
token_count_probe_metadata=main_model_config.get("token_count_probe_metadata")))
model_list.append(
ModelConfig(cite_name="sub_model",
api_key=main_model_config.get("api_key", ""),
Expand All @@ -907,7 +952,19 @@ async def create_model_config_list(tenant_id):
model_factory=main_model_config.get("model_factory"),
timeout_seconds=main_model_config.get("timeout_seconds"),
concurrency_limit=main_model_config.get("concurrency_limit"),
prompt_cache=main_prompt_cache))
prompt_cache=main_prompt_cache,
max_output_tokens=main_model_config.get("max_output_tokens"),
max_tokens=main_model_config.get("max_tokens"),
context_window_tokens=main_model_config.get("context_window_tokens"),
max_input_tokens=main_model_config.get("max_input_tokens"),
default_output_reserve_tokens=main_model_config.get("default_output_reserve_tokens"),
tokenizer_family=main_model_config.get("tokenizer_family"),
capacity_source=main_model_config.get("capacity_source"),
capability_profile_version=main_model_config.get("capability_profile_version"),
canonical_model_id=main_model_config.get("canonical_model_id"),
model_identity_metadata=main_model_config.get("model_identity_metadata"),
tokenizer_match_metadata=main_model_config.get("tokenizer_match_metadata"),
token_count_probe_metadata=main_model_config.get("token_count_probe_metadata")))

return model_list

Expand Down
Loading