Skip to content

Commit b7bbb84

Browse files
committed
fix: Ensure endpoints deploy command uses the deploy id and revision id when provided
1 parent d39b5a3 commit b7bbb84

7 files changed

Lines changed: 232 additions & 65 deletions

File tree

src/together/lib/cli/api/beta/endpoints/_utils/_resolve_model.py

Lines changed: 109 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from __future__ import annotations
22

33
import re
4+
from typing import NamedTuple
45

56
from together import APIError
67
from together.types.beta import Model, Endpoint
@@ -18,9 +19,12 @@
1819
# Logic for resolving a model + config from a user input string
1920
#
2021
# 1. Raw model id (e.g. ml_...)
21-
# → GET /configs?referenceModelId=... for the config and model path
22+
# → retrieve from --project when possible; config via baseModelId / id
23+
# → else GET /configs?referenceModelId=... (public / reference models)
2224
# 2. Full model path (projects/.../models/...)
23-
# → parse the model id, then same as (1)
25+
# → retrieve that model (keep it as the deploy target), resolve config via
26+
# baseModelId (or the model id when it is itself a reference model).
27+
# Preserve an optional /revisions/... pin on the deploy path.
2428
# 3. Named model (prefix/model-name)
2529
# a. prefix == project slug → list private models by name, resolve config via
2630
# baseModelId (path 1), but deploy the custom model path
@@ -31,7 +35,13 @@
3135
# - one profile + no --config → use that profile's config
3236
# - --config given → use the profile whose config matches
3337
# After either path, re-validate against the user's --config when provided.
34-
MODEL_PATH_RE = re.compile(r"^projects/([^/]+)/models/([^/]+)(?:/revisions/[^/]+)?$")
38+
MODEL_PATH_RE = re.compile(r"^projects/([^/]+)/models/([^/]+)(?:/revisions/([^/]+))?$")
39+
40+
41+
class ResolvedModelAndConfig(NamedTuple):
42+
model: Model
43+
config: Config
44+
revision_id: str | None = None
3545

3646

3747
async def resolve_model(
@@ -40,24 +50,48 @@ async def resolve_model(
4050
*,
4151
config_id: str | None = None,
4252
) -> Model:
43-
model, _config = await resolve_model_and_config(config, model_input, config_id=config_id)
44-
return model
53+
resolved = await resolve_model_and_config(config, model_input, config_id=config_id)
54+
return resolved.model
4555

4656

4757
async def resolve_model_and_config(
4858
config: CLIConfigParameter,
4959
model_input: str,
5060
*,
5161
config_id: str | None = None,
52-
) -> tuple[Model, Config]:
62+
) -> ResolvedModelAndConfig:
5363
"""Resolve a deployable model and the config revision to pair with it."""
54-
# 1 / 2. Full model path or raw model id → configs list by referenceModelId
64+
# 2. Full model path → keep the user's model; config from its base/reference.
5565
path_match = MODEL_PATH_RE.match(model_input)
5666
if path_match:
57-
_project_id, model_id = path_match.groups()
58-
return await _resolve_via_configs(config, model_id, config_id=config_id, model_input=model_input)
67+
project_id, model_id, revision_id = path_match.group(1), path_match.group(2), path_match.group(3)
68+
return await _resolve_explicit_model(
69+
config,
70+
model_id=model_id,
71+
project_id=project_id,
72+
config_id=config_id,
73+
model_input=model_input,
74+
revision_id=revision_id,
75+
)
5976

77+
# 1. Raw model id
6078
if "/" not in model_input:
79+
if config.project_id:
80+
try:
81+
model = await config.client.beta.models.retrieve(id=model_input, project_id=config.project_id)
82+
except APIError as e:
83+
if "not found" not in e.message.lower():
84+
raise
85+
else:
86+
reference_model_id = model.base_model_id or model.id
87+
assert reference_model_id is not None
88+
return await _resolve_config_for_model(
89+
config,
90+
model,
91+
reference_model_id=reference_model_id,
92+
config_id=config_id,
93+
model_input=model_input,
94+
)
6195
return await _resolve_via_configs(config, model_input, config_id=config_id, model_input=model_input)
6296

6397
# 3. Named model (prefix/model-name)
@@ -70,32 +104,84 @@ async def resolve_model_and_config(
70104
reference_model_id = model.base_model_id or model.id
71105
assert reference_model_id is not None
72106
# Config comes from the base/reference model; deploy path stays the custom model.
73-
_base_model, selected_config = await _resolve_via_configs(
107+
return await _resolve_config_for_model(
74108
config,
75-
reference_model_id,
109+
model,
110+
reference_model_id=reference_model_id,
76111
config_id=config_id,
77112
model_input=model_input,
78113
)
79-
return model, selected_config
80114

81115
return await _resolve_public_model_and_config(config, model_input, config_id=config_id)
82116

83117

118+
async def _resolve_explicit_model(
119+
config: CLIConfigParameter,
120+
*,
121+
model_id: str,
122+
project_id: str,
123+
config_id: str | None,
124+
model_input: str,
125+
revision_id: str | None = None,
126+
) -> ResolvedModelAndConfig:
127+
"""Load the user-specified model and pair it with a compatible config."""
128+
try:
129+
model = await config.client.beta.models.retrieve(id=model_id, project_id=project_id)
130+
except APIError as e:
131+
if "not found" in e.message.lower():
132+
raise ValueError(f"Model {model_input} not found.") from None
133+
raise
134+
135+
reference_model_id = model.base_model_id or model.id
136+
assert reference_model_id is not None
137+
return await _resolve_config_for_model(
138+
config,
139+
model,
140+
reference_model_id=reference_model_id,
141+
config_id=config_id,
142+
model_input=model_input,
143+
revision_id=revision_id,
144+
)
145+
146+
147+
async def _resolve_config_for_model(
148+
config: CLIConfigParameter,
149+
model: Model,
150+
*,
151+
reference_model_id: str,
152+
config_id: str | None,
153+
model_input: str,
154+
revision_id: str | None = None,
155+
) -> ResolvedModelAndConfig:
156+
selected = resolve_config(
157+
await resolve_configs(config, reference_model_id),
158+
config_id,
159+
model=model_input,
160+
)
161+
selected = validate_requested_config(selected, config_id, model=model_input)
162+
return ResolvedModelAndConfig(model=model, config=selected, revision_id=revision_id)
163+
164+
84165
async def _resolve_via_configs(
85166
config: CLIConfigParameter,
86167
reference_model_id: str,
87168
*,
88169
config_id: str | None,
89170
model_input: str,
90-
) -> tuple[Model, Config]:
171+
) -> ResolvedModelAndConfig:
172+
"""Resolve a public/reference model id through the configs API.
173+
174+
The deploy target is the config's reference model — correct when the user
175+
passed a bare reference-model id that is not retrievable under --project.
176+
"""
91177
selected = resolve_config(
92178
await resolve_configs(config, reference_model_id),
93179
config_id,
94180
model=model_input,
95181
)
96182
selected = validate_requested_config(selected, config_id, model=model_input)
97183
model = await _retrieve_model_from_reference(config, selected, model_input=model_input)
98-
return model, selected
184+
return ResolvedModelAndConfig(model=model, config=selected)
99185

100186

101187
async def _retrieve_model_from_reference(
@@ -108,7 +194,7 @@ async def _retrieve_model_from_reference(
108194
path = selected.reference_model or ""
109195
match = MODEL_PATH_RE.match(path)
110196
if match:
111-
project_id, model_id = match.groups()
197+
project_id, model_id = match.group(1), match.group(2)
112198
elif selected.reference_model_id and selected.project_id:
113199
project_id, model_id = selected.project_id, selected.reference_model_id
114200
else:
@@ -236,7 +322,7 @@ async def _resolve_public_model_and_config(
236322
model_input: str,
237323
*,
238324
config_id: str | None = None,
239-
) -> tuple[Model, Config]:
325+
) -> ResolvedModelAndConfig:
240326
supported_models = await config.client.beta.models.list_supported(search=model_input)
241327
if not supported_models.data:
242328
raise ValueError(f"Model {model_input} not found.")
@@ -255,19 +341,22 @@ async def _resolve_public_model_and_config(
255341
match = MODEL_PATH_RE.match(profile.model or "")
256342
if not match:
257343
raise ValueError(f"Invalid model path: {profile.model}")
258-
project_id, model_id = match.groups()
344+
project_id, model_id, revision_id = match.group(1), match.group(2), match.group(3)
259345

260346
selected_config = validate_requested_config(
261347
config_from_profile(profile),
262348
config_id,
263349
model=model_input,
264350
)
265351
model = Model.construct(id=model_id, projectId=project_id, name=public_model.name or model_id)
266-
return model, selected_config
352+
return ResolvedModelAndConfig(model=model, config=selected_config, revision_id=revision_id)
267353

268354

269-
def construct_model_path(model: Model) -> str:
270-
return f"projects/{model.project_id}/models/{model.id}"
355+
def construct_model_path(model: Model, revision_id: str | None = None) -> str:
356+
path = f"projects/{model.project_id}/models/{model.id}"
357+
if revision_id:
358+
return f"{path}/revisions/{revision_id}"
359+
return path
271360

272361

273362
async def resolve_endpoint(config: CLIConfigParameter, endpoint_id_or_name: str) -> Endpoint:

src/together/lib/cli/api/beta/endpoints/ab.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -74,7 +74,8 @@ async def ab(
7474
)
7575
verify_control_receiving_traffic(endpoint, control)
7676

77-
resolved_model, config_value = await resolve_model_and_config(config, model, config_id=config_id)
77+
resolved = await resolve_model_and_config(config, model, config_id=config_id)
78+
resolved_model, config_value = resolved.model, resolved.config
7879

7980
autoscaling = build_autoscaling(
8081
min_replicas=1,
@@ -103,7 +104,7 @@ async def ab(
103104
endpoint_id=endpoint.id,
104105
enable_lora=enable_lora,
105106
name=name,
106-
model=construct_model_path(resolved_model),
107+
model=construct_model_path(resolved_model, resolved.revision_id),
107108
config=construct_config_path(config_value),
108109
autoscaling=autoscaling,
109110
),

src/together/lib/cli/api/beta/endpoints/deploy.py

Lines changed: 13 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -174,7 +174,11 @@ async def deploy(
174174
config: CLIConfigParameter,
175175
) -> None:
176176
"""Create a deployment on a new or existing dedicated inference endpoint."""
177-
resolved_model, config_value = await resolve_model_and_config(config, model, config_id=config_id)
177+
resolved = await resolve_model_and_config(config, model, config_id=config_id)
178+
resolved_model, config_value = resolved.model, resolved.config
179+
# Prefer revision pin from a fully-qualified model path; fall back to the
180+
# deprecated --model-revision flag.
181+
resolved_revision = resolved.revision_id or model_revision
178182

179183
autoscaling = build_autoscaling(
180184
min_replicas=min_replicas,
@@ -200,16 +204,18 @@ async def deploy(
200204
else:
201205
placement_value = placement.to_json()
202206

207+
model_path = construct_model_path(resolved_model, resolved_revision)
208+
203209
if not config.json:
204210
_print_deployment_preview(
205211
endpoint=endpoint_name_or_id,
206212
deployment_name=deployment_name,
207213
model=resolved_model,
214+
model_path=model_path,
208215
config_value=config_value,
209216
autoscaling=autoscaling,
210217
placement=placement_value,
211218
enable_lora=enable_lora,
212-
model_revision=model_revision,
213219
traffic_weight=traffic_weight,
214220
)
215221
await assert_explicit_project_id(config)
@@ -222,11 +228,12 @@ async def deploy(
222228
config.client.beta.endpoints.deployments.create(
223229
endpoint.id,
224230
name=deployment_name,
225-
model=construct_model_path(resolved_model),
231+
model=model_path,
226232
config=construct_config_path(config_value),
227233
autoscaling=autoscaling,
228234
enable_lora=enable_lora if enable_lora is not None else omit,
229-
model_revision_id=model_revision or omit,
235+
# Revision is already embedded in model_path when present.
236+
model_revision_id=omit,
230237
placement=placement_value or omit,
231238
),
232239
)
@@ -267,11 +274,11 @@ def _print_deployment_preview(
267274
endpoint: str,
268275
deployment_name: str,
269276
model: Model,
277+
model_path: str,
270278
config_value: Config,
271279
autoscaling: DeploymentAutoscalingParam,
272280
placement: Placement | None,
273281
enable_lora: bool | None,
274-
model_revision: str | None,
275282
traffic_weight: float | None,
276283
) -> None:
277284
table = Table(expand=True, show_header=False, show_edge=False, show_lines=False, box=None, pad_edge=False)
@@ -302,9 +309,6 @@ def add_row(flag: str, value: str) -> None:
302309
if percentile := metric.get("percentile"):
303310
add_row("--scaling-percentile", percentile)
304311

305-
if model_revision:
306-
add_row("--model-revision", model_revision)
307-
308312
if placement is not None:
309313
if "profile" in placement:
310314
add_row("--placement", placement["profile"]) # type: ignore[typeddict-item]
@@ -324,7 +328,7 @@ def add_row(flag: str, value: str) -> None:
324328
add_row("--enable-lora", "true" if enable_lora else "false")
325329
if traffic_weight is not None:
326330
add_row("--traffic-weight", str(traffic_weight))
327-
add_row("--model", model.name)
331+
add_row("--model", f"{model.name} ({model_path})")
328332
add_row("--config", config_value.id) # type: ignore
329333

330334
table.add_row("\n".join(args))

src/together/lib/cli/api/beta/endpoints/shadow.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -92,7 +92,8 @@ async def shadow(
9292
rate, target_qps = await resolve_rate_or_target_qps(rate, target_qps, config=config)
9393

9494
endpoint_id = (await resolve_endpoint(config, endpoint_id_or_name)).id
95-
resolved_model, config_value = await resolve_model_and_config(config, model, config_id=config_id)
95+
resolved = await resolve_model_and_config(config, model, config_id=config_id)
96+
resolved_model, config_value = resolved.model, resolved.config
9697

9798
autoscaling = build_autoscaling(
9899
min_replicas=1,
@@ -122,7 +123,7 @@ async def shadow(
122123
endpoint_id=endpoint_id,
123124
name=name,
124125
enable_lora=enable_lora,
125-
model=construct_model_path(resolved_model),
126+
model=construct_model_path(resolved_model, resolved.revision_id),
126127
config=construct_config_path(config_value),
127128
autoscaling=autoscaling,
128129
),

src/together/lib/cli/api/beta/models/create.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@ async def preprompt(self, config: CLIConfig) -> None:
3030
for profile in model.deployment_profiles:
3131
match = MODEL_PATH_RE.match(profile.model)
3232
if match:
33-
_, model_id = match.groups()
33+
model_id = match.group(2)
3434
self.choices.append((f"{model.name} ({profile.quantization})", model_id))
3535

3636

tests/cli/test_beta_endpoints_shadow.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -404,10 +404,12 @@ def test_shadow_resolves_single_config(
404404

405405
assert result.exit_code == 0, result.output
406406
url = str(cast(Call, configs_route.calls[0]).request.url)
407-
assert "referenceModelId=ml_1" in url
407+
# Configs are resolved via the model's baseModelId; deploy target stays ml_1.
408+
assert "referenceModelId=ml_base" in url
408409

409410
deployment_body = json.loads(cast(Call, create_deployment_route.calls[0]).request.content.decode())
410411
assert deployment_body["config"] == "projects/proj/configs/cr_1"
412+
assert deployment_body["model"] == "projects/proj/models/ml_1"
411413

412414
@pytest.mark.respx(base_url=base_url)
413415
def test_shadow_errors_on_multiple_configs(
@@ -416,6 +418,7 @@ def test_shadow_errors_on_multiple_configs(
416418
cli_runner: CliRunner,
417419
) -> None:
418420
_mock_endpoint(respx_mock)
421+
respx_mock.get("/projects/proj/models/ml_1").mock(return_value=httpx.Response(200, json=_model_body()))
419422
respx_mock.get("/projects/proj/configs").mock(
420423
return_value=httpx.Response(
421424
200,

0 commit comments

Comments
 (0)