diff --git a/activity/llm_inference.go b/activity/llm_inference.go index d7ce112..f9a10de 100644 --- a/activity/llm_inference.go +++ b/activity/llm_inference.go @@ -3,6 +3,7 @@ package activity import ( "context" "fmt" + "os" "github.com/rockliang/poimen/workflows/activity/llm" "github.com/rockliang/poimen/workflows/pkg/types" @@ -44,11 +45,17 @@ func LLMInferenceActivity(ctx context.Context, in LLMInferenceInput) (LLMInferen return output, fmt.Errorf("failed to create LLM client: %w", err) } + // Use provided auth token, or fallback to environment variable + authToken := in.AuthToken + if authToken == "" { + authToken = os.Getenv("LLM_AUTH_TOKEN") + } + response, err := client.CreateMessage(ctx, llm.MessageInput{ Model: types.ModelSpec{ModelID: in.Model}, SystemPrompt: in.SystemPrompt, Messages: []llm.MessageParam{{Role: "user", Content: in.UserPrompt}}, - AuthToken: in.AuthToken, + AuthToken: authToken, }) if err != nil { output.ErrorMessage = err.Error() @@ -94,11 +101,18 @@ func LLMBatchInferenceActivity(ctx context.Context, in LLMBatchInferenceInput) ( return output, fmt.Errorf("failed to create LLM client: %w", err) } + // Use provided auth token, or fallback to environment variable + authToken := in.AuthToken + if authToken == "" { + authToken = os.Getenv("LLM_AUTH_TOKEN") + } + for i, prompt := range in.Prompts { response, err := client.CreateMessage(ctx, llm.MessageInput{ Model: types.ModelSpec{ModelID: in.Model}, SystemPrompt: in.SystemPrompt, Messages: []llm.MessageParam{{Role: "user", Content: prompt}}, + AuthToken: authToken, }) if err != nil { output.Errors = append(output.Errors, fmt.Sprintf("prompt %d: %v", i, err)) diff --git a/poimen-worker b/poimen-worker new file mode 100755 index 0000000..17509c9 Binary files /dev/null and b/poimen-worker differ