Skip to content

Commit 4d8ec08

Browse files
authored
refactor: consolidate dynamic provider config parsing (#4985)
1 parent 668d7fa commit 4d8ec08

10 files changed

Lines changed: 581 additions & 170 deletions

File tree

src/llama_stack/cli/stack/_list_deps.py

Lines changed: 7 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -12,10 +12,9 @@
1212
from termcolor import cprint
1313

1414
from llama_stack.core.build import get_provider_dependencies
15-
from llama_stack.core.datatypes import Provider, StackConfig
16-
from llama_stack.core.distribution import get_provider_registry
15+
from llama_stack.core.datatypes import StackConfig
16+
from llama_stack.core.stack import run_config_from_dynamic_config_spec
1717
from llama_stack.log import get_logger
18-
from llama_stack_api import Api
1918

2019
TEMPLATES_PATH = Path(__file__).parent.parent.parent / "templates"
2120

@@ -91,39 +90,11 @@ def run_stack_list_deps_command(args: argparse.Namespace) -> None:
9190
)
9291
sys.exit(1)
9392
elif args.providers:
94-
provider_list: dict[str, list[Provider]] = dict()
95-
for api_provider in args.providers.split(","):
96-
if "=" not in api_provider:
97-
cprint(
98-
"Could not parse `--providers`. Please ensure the list is in the format api1=provider1,api2=provider2",
99-
color="red",
100-
file=sys.stderr,
101-
)
102-
sys.exit(1)
103-
api, provider_type = api_provider.split("=")
104-
providers_for_api = get_provider_registry().get(Api(api), None)
105-
if providers_for_api is None:
106-
cprint(
107-
f"{api} is not a valid API.",
108-
color="red",
109-
file=sys.stderr,
110-
)
111-
sys.exit(1)
112-
if provider_type in providers_for_api:
113-
provider = Provider(
114-
provider_type=provider_type,
115-
provider_id=provider_type.split("::")[1],
116-
module=None,
117-
)
118-
provider_list.setdefault(api, []).append(provider)
119-
else:
120-
cprint(
121-
f"{provider_type} is not a valid provider for the {api} API.",
122-
color="red",
123-
file=sys.stderr,
124-
)
125-
sys.exit(1)
126-
config = StackConfig(providers=provider_list, distro_name="providers-run")
93+
try:
94+
config = run_config_from_dynamic_config_spec(args.providers)
95+
except ValueError as e:
96+
cprint(str(e), color="red", file=sys.stderr)
97+
sys.exit(1)
12798

12899
normal_deps, special_deps, external_provider_dependencies = get_provider_dependencies(config)
129100
normal_deps += SERVER_DEPENDENCIES

src/llama_stack/cli/stack/run.py

Lines changed: 13 additions & 95 deletions
Original file line numberDiff line numberDiff line change
@@ -16,21 +16,10 @@
1616
from termcolor import cprint
1717

1818
from llama_stack.cli.subcommand import Subcommand
19-
from llama_stack.core.datatypes import Api, Provider, StackConfig
20-
from llama_stack.core.distribution import get_provider_registry
21-
from llama_stack.core.stack import cast_distro_name_to_string, replace_env_vars
22-
from llama_stack.core.storage.datatypes import (
23-
InferenceStoreReference,
24-
KVStoreReference,
25-
ServerStoresConfig,
26-
SqliteKVStoreConfig,
27-
SqliteSqlStoreConfig,
28-
SqlStoreReference,
29-
StorageConfig,
30-
)
19+
from llama_stack.core.datatypes import StackConfig
20+
from llama_stack.core.stack import cast_distro_name_to_string, replace_env_vars, run_config_from_dynamic_config_spec
3121
from llama_stack.core.utils.config_dirs import DISTRIBS_BASE_DIR
3222
from llama_stack.core.utils.config_resolution import resolve_config_or_distro
33-
from llama_stack.core.utils.dynamic import instantiate_class_type
3423
from llama_stack.log import LoggingConfig, get_logger
3524

3625
REPO_ROOT = Path(__file__).parent.parent.parent.parent
@@ -92,50 +81,20 @@ def _run_stack_run_cmd(self, args: argparse.Namespace) -> None:
9281
except ValueError as e:
9382
self.parser.error(str(e))
9483
elif args.providers:
95-
provider_list: dict[str, list[Provider]] = dict()
96-
for api_provider in args.providers.split(","):
97-
if "=" not in api_provider:
98-
cprint(
99-
"Could not parse `--providers`. Please ensure the list is in the format api1=provider1,api2=provider2",
100-
color="red",
101-
file=sys.stderr,
102-
)
103-
sys.exit(1)
104-
api, provider_type = api_provider.split("=")
105-
providers_for_api = get_provider_registry().get(Api(api), None)
106-
if providers_for_api is None:
107-
cprint(
108-
f"{api} is not a valid API.",
109-
color="red",
110-
file=sys.stderr,
111-
)
112-
sys.exit(1)
113-
if provider_type in providers_for_api:
114-
config_type = instantiate_class_type(providers_for_api[provider_type].config_class)
115-
if config_type is not None and hasattr(config_type, "sample_run_config"):
116-
config = config_type.sample_run_config(__distro_dir__="~/.llama/distributions/providers-run")
117-
else:
118-
config = {}
119-
provider = Provider(
120-
provider_type=provider_type,
121-
config=config,
122-
provider_id=provider_type.split("::")[1],
123-
)
124-
provider_list.setdefault(api, []).append(provider)
125-
else:
126-
cprint(
127-
f"{provider} is not a valid provider for the {api} API.",
128-
color="red",
129-
file=sys.stderr,
130-
)
131-
sys.exit(1)
132-
run_config = self._generate_run_config_from_providers(providers=provider_list)
84+
distro_dir = DISTRIBS_BASE_DIR / "providers-run"
85+
os.makedirs(distro_dir, exist_ok=True)
86+
try:
87+
run_config = run_config_from_dynamic_config_spec(
88+
dynamic_config_spec=args.providers,
89+
distro_dir=distro_dir,
90+
distro_name="providers-run",
91+
)
92+
except ValueError as e:
93+
cprint(str(e), color="red", file=sys.stderr)
94+
sys.exit(1)
13395
config_dict = run_config.model_dump(mode="json")
13496

135-
# Write config to disk in providers-run directory
136-
distro_dir = DISTRIBS_BASE_DIR / "providers-run"
13797
config_file = distro_dir / "config.yaml"
138-
13998
logger.info(f"Writing generated config to: {config_file}")
14099
with open(config_file, "w") as f:
141100
yaml.dump(config_dict, f, default_flow_style=False, sort_keys=False)
@@ -261,44 +220,3 @@ def _start_ui_development_server(self, stack_server_port: int):
261220
)
262221
except Exception as e:
263222
logger.error(f"Failed to start UI development server in {ui_dir}: {e}")
264-
265-
def _generate_run_config_from_providers(self, providers: dict[str, list[Provider]]):
266-
apis = list(providers.keys())
267-
distro_dir = DISTRIBS_BASE_DIR / "providers-run"
268-
# need somewhere to put the storage.
269-
os.makedirs(distro_dir, exist_ok=True)
270-
storage = StorageConfig(
271-
backends={
272-
"kv_default": SqliteKVStoreConfig(
273-
db_path=f"${{env.SQLITE_STORE_DIR:={distro_dir}}}/kvstore.db",
274-
),
275-
"sql_default": SqliteSqlStoreConfig(
276-
db_path=f"${{env.SQLITE_STORE_DIR:={distro_dir}}}/sql_store.db",
277-
),
278-
},
279-
stores=ServerStoresConfig(
280-
metadata=KVStoreReference(
281-
backend="kv_default",
282-
namespace="registry",
283-
),
284-
inference=InferenceStoreReference(
285-
backend="sql_default",
286-
table_name="inference_store",
287-
),
288-
conversations=SqlStoreReference(
289-
backend="sql_default",
290-
table_name="openai_conversations",
291-
),
292-
prompts=KVStoreReference(
293-
backend="kv_default",
294-
namespace="prompts",
295-
),
296-
),
297-
)
298-
299-
return StackConfig(
300-
distro_name="providers-run",
301-
apis=apis,
302-
providers=providers,
303-
storage=storage,
304-
)

src/llama_stack/core/stack.py

Lines changed: 31 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
import os
1111
import re
1212
import tempfile
13+
from pathlib import Path
1314
from typing import Any, get_type_hints
1415

1516
import yaml
@@ -728,26 +729,38 @@ def get_stack_run_config_from_distro(distro: str) -> StackConfig:
728729
return StackConfig(**replace_env_vars(run_config))
729730

730731

731-
def run_config_from_adhoc_config_spec(
732-
adhoc_config_spec: str, provider_registry: ProviderRegistry | None = None
732+
def run_config_from_dynamic_config_spec(
733+
dynamic_config_spec: str,
734+
provider_registry: ProviderRegistry | None = None,
735+
distro_dir: Path | None = None,
736+
distro_name: str = "dynamic-distro",
733737
) -> StackConfig:
734738
"""
735-
Create an adhoc distribution from a list of API providers.
739+
Create a dynamic distribution from a list of API providers.
736740
737741
The list should be of the form "api=provider", e.g. "inference=fireworks". If you have
738742
multiple pairs, separate them with commas or semicolons, e.g. "inference=fireworks,safety=llama-guard,agents=meta-reference"
739743
"""
740744

741-
api_providers = adhoc_config_spec.replace(";", ",").split(",")
742-
provider_registry = provider_registry or get_provider_registry()
745+
api_providers = dynamic_config_spec.replace(";", ",").split(",")
746+
provider_registry = get_provider_registry() if provider_registry is None else provider_registry
743747

744-
distro_dir = tempfile.mkdtemp()
748+
distro_dir = distro_dir or Path(tempfile.mkdtemp())
745749
provider_configs_by_api = {}
746750
for api_provider in api_providers:
751+
if "=" not in api_provider:
752+
raise ValueError(
753+
f"Failed to parse provider spec '{api_provider}'. Expected format: api=provider (e.g. inference=fireworks)"
754+
)
747755
api_str, provider = api_provider.split("=")
748-
api = Api(api_str)
756+
try:
757+
api = Api(api_str)
758+
except ValueError:
759+
raise ValueError(f"Failed to parse provider spec: '{api_str}' is not a valid API") from None
749760

750-
providers_by_type = provider_registry[api]
761+
providers_by_type = provider_registry.get(api)
762+
if providers_by_type is None:
763+
raise ValueError(f"Failed to find providers for API '{api_str}'")
751764
provider_spec = providers_by_type.get(provider)
752765
if not provider_spec:
753766
provider_spec = providers_by_type.get(f"inline::{provider}")
@@ -761,23 +774,27 @@ def run_config_from_adhoc_config_spec(
761774

762775
# call method "sample_run_config" on the provider spec config class
763776
provider_config_type = instantiate_class_type(provider_spec.config_class)
764-
provider_config = replace_env_vars(provider_config_type.sample_run_config(__distro_dir__=distro_dir))
777+
provider_config = replace_env_vars(provider_config_type.sample_run_config(__distro_dir__=str(distro_dir)))
765778

766-
provider_configs_by_api[api_str] = [
779+
provider_configs_by_api.setdefault(api_str, []).append(
767780
Provider(
768781
provider_id=provider_spec.provider_type.split("::")[-1],
769782
provider_type=provider_spec.provider_type,
770783
config=provider_config,
771784
)
772-
]
785+
)
773786
config = StackConfig(
774-
distro_name="distro-test",
787+
distro_name=distro_name,
775788
apis=list(provider_configs_by_api.keys()),
776789
providers=provider_configs_by_api,
777790
storage=StorageConfig(
778791
backends={
779-
"kv_default": SqliteKVStoreConfig(db_path=f"{distro_dir}/kvstore.db"),
780-
"sql_default": SqliteSqlStoreConfig(db_path=f"{distro_dir}/sql_store.db"),
792+
"kv_default": SqliteKVStoreConfig(
793+
db_path=f"${{env.SQLITE_STORE_DIR:={distro_dir}}}/kvstore.db",
794+
),
795+
"sql_default": SqliteSqlStoreConfig(
796+
db_path=f"${{env.SQLITE_STORE_DIR:={distro_dir}}}/sql_store.db",
797+
),
781798
},
782799
stores=ServerStoresConfig(
783800
metadata=KVStoreReference(backend="kv_default", namespace="registry"),

tests/integration/conftest.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@
1515
import pytest
1616
from dotenv import load_dotenv
1717

18-
from llama_stack.core.stack import run_config_from_adhoc_config_spec
18+
from llama_stack.core.stack import run_config_from_dynamic_config_spec
1919
from llama_stack.log import get_logger
2020
from llama_stack.testing.api_recorder import patch_httpx_for_test_id
2121

@@ -148,7 +148,7 @@ def pytest_configure(config):
148148
if getattr(config.option, "embedding_model", None) is None:
149149
stack_config = config.getoption("--stack-config", default=None)
150150
if stack_config and "=" in stack_config:
151-
run_config = run_config_from_adhoc_config_spec(stack_config)
151+
run_config = run_config_from_dynamic_config_spec(stack_config)
152152
inference_providers = run_config.providers.get("inference", [])
153153
if any("sentence-transformers" in p.provider_type for p in inference_providers):
154154
config.option.embedding_model = "sentence-transformers/nomic-ai/nomic-embed-text-v1.5"
@@ -162,7 +162,7 @@ def pytest_addoption(parser):
162162
a 'pointer' to the stack. this can be either be:
163163
(a) a template name like `starter`, or
164164
(b) a path to a config.yaml file, or
165-
(c) an adhoc config spec, e.g. `inference=fireworks,safety=llama-guard,agents=meta-reference`, or
165+
(c) a dynamic config spec, e.g. `inference=fireworks,safety=llama-guard,agents=meta-reference`, or
166166
(d) a server config like `server:ci-tests`, or
167167
(e) a docker config like `docker:ci-tests` (builds and runs container)
168168
"""
@@ -256,7 +256,7 @@ def pytest_generate_tests(metafunc):
256256
config_str = metafunc.config.getoption("--stack-config", default=None) or os.environ.get("LLAMA_STACK_CONFIG")
257257
providers = None
258258
if config_str and "=" in config_str:
259-
run_config = run_config_from_adhoc_config_spec(config_str)
259+
run_config = run_config_from_dynamic_config_spec(config_str)
260260
providers = [p.provider_id for p in run_config.providers.get("vector_io", [])]
261261
if providers is None:
262262
inference_mode = os.environ.get("LLAMA_STACK_TEST_INFERENCE_MODE")

tests/integration/fixtures/common.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,7 @@
2727

2828
from llama_stack.core.datatypes import QualifiedModel, VectorStoresConfig
2929
from llama_stack.core.library_client import LlamaStackAsLibraryClient
30-
from llama_stack.core.stack import run_config_from_adhoc_config_spec
30+
from llama_stack.core.stack import run_config_from_dynamic_config_spec
3131
from llama_stack.core.utils.config_resolution import resolve_config_or_distro
3232
from llama_stack.env import get_env_or_fail
3333

@@ -321,7 +321,7 @@ def instantiate_llama_stack_client(session):
321321
pass
322322

323323
if "=" in config:
324-
run_config = run_config_from_adhoc_config_spec(config)
324+
run_config = run_config_from_dynamic_config_spec(config)
325325

326326
# --stack-config bypasses template so need this to set default embedding model
327327
if "vector_io" in config and "inference" in config:

tests/unit/cli/test_stack_config.py

Lines changed: 6 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -229,26 +229,13 @@ def test_parse_and_maybe_upgrade_config_preserves_custom_external_providers_dir(
229229

230230

231231
def test_generate_run_config_from_providers():
232-
"""Test that _generate_run_config_from_providers creates a valid config"""
233-
import argparse
234-
235-
from llama_stack.cli.stack.run import StackRun
236-
from llama_stack.core.datatypes import Provider
232+
"""Test that run_config_from_dynamic_config_spec creates a valid config for the providers-run distro"""
233+
from llama_stack.core.stack import run_config_from_dynamic_config_spec
237234

238-
parser = argparse.ArgumentParser()
239-
subparsers = parser.add_subparsers()
240-
stack_run = StackRun(subparsers)
241-
242-
providers = {
243-
"inference": [
244-
Provider(
245-
provider_type="inline::meta-reference",
246-
provider_id="meta-reference",
247-
)
248-
]
249-
}
250-
251-
config = stack_run._generate_run_config_from_providers(providers=providers)
235+
config = run_config_from_dynamic_config_spec(
236+
"inference=remote::openai",
237+
distro_name="providers-run",
238+
)
252239
config_dict = config.model_dump(mode="json")
253240

254241
# Verify basic structure

0 commit comments

Comments
 (0)