Files
notarius/internal/cli/catalog.go

218 lines
8.2 KiB
Go

package cli
import (
"context"
"fmt"
"gitea.maximumdirect.net/eric/notarius/internal/core/artifacts"
"gitea.maximumdirect.net/eric/notarius/internal/core/config"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/llm"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
"gitea.maximumdirect.net/eric/notarius/internal/modules/chunk/dnd/scenes"
"gitea.maximumdirect.net/eric/notarius/internal/modules/chunk/generic"
"gitea.maximumdirect.net/eric/notarius/internal/modules/extract/dnd/spells"
"gitea.maximumdirect.net/eric/notarius/internal/modules/input/seriatim"
"gitea.maximumdirect.net/eric/notarius/internal/modules/merge/appendorder"
"gitea.maximumdirect.net/eric/notarius/internal/modules/normalize/noop"
jsonoutput "gitea.maximumdirect.net/eric/notarius/internal/modules/output/json"
spellshape "gitea.maximumdirect.net/eric/notarius/internal/validators/extract/dnd/spells/shape"
spellsourcerefs "gitea.maximumdirect.net/eric/notarius/internal/validators/extract/dnd/spells/source_refs"
spellrelatedness "gitea.maximumdirect.net/eric/notarius/internal/validators/extract/dnd/spells/source_relatedness"
alwaysaccept "gitea.maximumdirect.net/eric/notarius/internal/validators/generic/always_accept"
alwaysreject "gitea.maximumdirect.net/eric/notarius/internal/validators/generic/always_reject"
validjson "gitea.maximumdirect.net/eric/notarius/internal/validators/generic/valid_json"
validjsonschema "gitea.maximumdirect.net/eric/notarius/internal/validators/generic/valid_json_schema"
)
func productionRegistries() (pipeline.Registries, error) {
registries := pipeline.Registries{
Inputs: pipeline.NewInputAdapterRegistry(),
Chunkers: pipeline.NewChunkerRegistry(),
Extractors: pipeline.NewExtractorRegistry(),
Mergers: pipeline.NewMergerRegistry(),
Normalizers: pipeline.NewNormalizerRegistry(),
Validators: pipeline.NewValidatorRegistry(),
ValidatorChains: pipeline.NewValidatorChainRegistry(),
Outputs: pipeline.NewOutputEncoderRegistry(),
}
if err := seriatim.Register(registries.Inputs); err != nil {
return pipeline.Registries{}, fmt.Errorf("register seriatim input: %w", err)
}
if err := generic.Register(registries.Chunkers); err != nil {
return pipeline.Registries{}, fmt.Errorf("register generic chunker: %w", err)
}
if err := scenes.Register(registries.Chunkers); err != nil {
return pipeline.Registries{}, fmt.Errorf("register dnd scenes chunker: %w", err)
}
if err := spells.Register(registries.Extractors); err != nil {
return pipeline.Registries{}, fmt.Errorf("register dnd spells extractor: %w", err)
}
if err := appendorder.Register(registries.Mergers); err != nil {
return pipeline.Registries{}, fmt.Errorf("register appendorder merger: %w", err)
}
if err := noop.Register(registries.Normalizers); err != nil {
return pipeline.Registries{}, fmt.Errorf("register noop normalizer: %w", err)
}
if err := registerProductionValidators(registries.Validators); err != nil {
return pipeline.Registries{}, err
}
if err := registerProductionValidatorChains(registries.ValidatorChains); err != nil {
return pipeline.Registries{}, err
}
if err := jsonoutput.Register(registries.Outputs); err != nil {
return pipeline.Registries{}, fmt.Errorf("register json output encoder: %w", err)
}
return registries, nil
}
func registerProductionValidators(registry *pipeline.ValidatorRegistry) error {
registrations := []struct {
name string
register func(*pipeline.ValidatorRegistry) error
}{
{name: "generic always accept validator", register: alwaysaccept.Register},
{name: "generic always reject validator", register: alwaysreject.Register},
{name: "generic valid json validator", register: validjson.Register},
{name: "generic valid json schema validator", register: validjsonschema.Register},
{name: "dnd spell shape validator", register: spellshape.Register},
{name: "dnd spell source references validator", register: spellsourcerefs.Register},
{name: "dnd spell source relatedness validator", register: spellrelatedness.Register},
}
for _, registration := range registrations {
if err := registration.register(registry); err != nil {
return fmt.Errorf("register %s: %w", registration.name, err)
}
}
return nil
}
func registerProductionValidatorChains(registry *pipeline.ValidatorChainRegistry) error {
if err := registry.Register(pipeline.ValidatorChainMapping{
Stage: pipeline.StageExtract,
Module: spells.Key,
Validators: []pipeline.ModuleBinding{
pipeline.Binding(validjson.Key),
pipeline.Binding(validjsonschema.Key),
pipeline.Binding(spellshape.Key),
pipeline.Binding(spellsourcerefs.Key),
pipeline.Binding(spellrelatedness.Key),
},
}); err != nil {
return fmt.Errorf("register dnd spells validator chain: %w", err)
}
return nil
}
func productionCatalog() (pipeline.ModuleCatalog, error) {
registries, err := productionRegistries()
if err != nil {
return pipeline.ModuleCatalog{}, err
}
return catalogFromRegistries(registries), nil
}
func productionPromptAssets() (*llm.AssetRegistry, error) {
registry := llm.NewAssetRegistry()
if err := scenes.RegisterPromptAssets(registry); err != nil {
return nil, fmt.Errorf("register dnd scenes prompt assets: %w", err)
}
if err := spells.RegisterPromptAssets(registry); err != nil {
return nil, fmt.Errorf("register dnd spells prompt assets: %w", err)
}
return registry, nil
}
func effectiveCatalog(opts Options) (pipeline.ModuleCatalog, error) {
if !isEmptyCatalog(opts.Catalog) {
return opts.Catalog, nil
}
if !isEmptyRegistries(opts.Registries) {
return catalogFromRegistries(opts.Registries), nil
}
return productionCatalog()
}
func effectiveRegistries(opts Options) (pipeline.Registries, error) {
if !isEmptyRegistries(opts.Registries) {
return opts.Registries, nil
}
if !isEmptyCatalog(opts.Catalog) {
return registriesFromCatalog(opts.Catalog), nil
}
return productionRegistries()
}
func catalogFromRegistries(registries pipeline.Registries) pipeline.ModuleCatalog {
return pipeline.ModuleCatalog{
Inputs: registries.Inputs,
Chunkers: registries.Chunkers,
Extractors: registries.Extractors,
Mergers: registries.Mergers,
Normalizers: registries.Normalizers,
Validators: registries.Validators,
ValidatorChains: registries.ValidatorChains,
Outputs: registries.Outputs,
}
}
func registriesFromCatalog(catalog pipeline.ModuleCatalog) pipeline.Registries {
return pipeline.Registries{
Inputs: catalog.Inputs,
Chunkers: catalog.Chunkers,
Extractors: catalog.Extractors,
Mergers: catalog.Mergers,
Normalizers: catalog.Normalizers,
Validators: catalog.Validators,
ValidatorChains: catalog.ValidatorChains,
Outputs: catalog.Outputs,
}
}
func isEmptyCatalog(catalog pipeline.ModuleCatalog) bool {
return catalog.Inputs == nil &&
catalog.Chunkers == nil &&
catalog.Extractors == nil &&
catalog.Mergers == nil &&
catalog.Normalizers == nil &&
catalog.Validators == nil &&
catalog.ValidatorChains == nil &&
catalog.Outputs == nil
}
func isEmptyRegistries(registries pipeline.Registries) bool {
return registries.Inputs == nil &&
registries.Chunkers == nil &&
registries.Extractors == nil &&
registries.Mergers == nil &&
registries.Normalizers == nil &&
registries.Validators == nil &&
registries.ValidatorChains == nil &&
registries.Outputs == nil
}
func productionLLMClientFactory(ctx context.Context, cfg config.Config, profileID string) (contracts.StructuredLLMClient, []artifacts.LLMProfileManifest, error) {
if err := ctx.Err(); err != nil {
return nil, nil, err
}
assets, err := productionPromptAssets()
if err != nil {
return nil, nil, err
}
recorder := llm.NewLLMProfileRecorder()
client, err := llm.NewScriptoriumClient(llm.ScriptoriumClientConfig{
ProfileDir: cfg.Scriptorium.ProfileDir,
ProfileFile: cfg.Scriptorium.ProfileFile,
Assets: assets,
Recorder: recorder,
})
if err != nil {
return nil, nil, fmt.Errorf("create Scriptorium-backed LLM client: %w", err)
}
scheduler, err := llm.NewScheduler(cfg.Concurrency.TotalLLM)
if err != nil {
return nil, nil, fmt.Errorf("create LLM scheduler: %w", err)
}
return llm.NewScheduledClient(client, scheduler), nil, nil
}