From 92f9ba1550dfa3d3baaab71e418ff18efa3052e5 Mon Sep 17 00:00:00 2001 From: gohumble Date: Tue, 30 Aug 2022 16:15:14 +0200 Subject: [PATCH] added dalle --- internal/dalle/client.go | 14 ++ internal/dalle/dalle.go | 11 ++ internal/dalle/httpclient.go | 259 ++++++++++++++++++++++++++++++++++ internal/telegram/generate.go | 50 ++++--- 4 files changed, 316 insertions(+), 18 deletions(-) create mode 100644 internal/dalle/client.go create mode 100644 internal/dalle/dalle.go create mode 100644 internal/dalle/httpclient.go diff --git a/internal/dalle/client.go b/internal/dalle/client.go new file mode 100644 index 0000000..42a7ca5 --- /dev/null +++ b/internal/dalle/client.go @@ -0,0 +1,14 @@ +package dalle + +import ( + "context" + "io" +) + +type Client interface { + Generate(ctx context.Context, prompt string) (*Task, error) + ListTasks(ctx context.Context, req *ListTasksRequest) (*ListTasksResponse, error) + GetTask(ctx context.Context, taskID string) (*Task, error) + Download(ctx context.Context, generationID string) (io.ReadCloser, error) + Share(ctx context.Context, generationID string) (string, error) +} diff --git a/internal/dalle/dalle.go b/internal/dalle/dalle.go new file mode 100644 index 0000000..7a96853 --- /dev/null +++ b/internal/dalle/dalle.go @@ -0,0 +1,11 @@ +package dalle + +const ( + StatusPending = "pending" + StatusRejected = "rejected" + StatusSucceeded = "succeeded" + + TaskTypeText2Im = "text2im" + + defaultBatchSize = 4 +) diff --git a/internal/dalle/httpclient.go b/internal/dalle/httpclient.go new file mode 100644 index 0000000..c0c4aff --- /dev/null +++ b/internal/dalle/httpclient.go @@ -0,0 +1,259 @@ +package dalle + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/url" + "time" +) + +const ( + libraryVersion = "1.0.0" + defaultUserAgent = "dalle/" + libraryVersion + baseURL = "https://labs.openai.com/api/labs" + + defaultHTTPClientTimeout = 15 * time.Second +) + +type option func(*HTTPClient) error + +func WithHTTPClient(httpClient *http.Client) option { + return func(c *HTTPClient) error { + c.httpClient = httpClient + + return nil + } +} + +func WithUserAgent(userAgent string) option { + return func(c *HTTPClient) error { + c.userAgent = userAgent + + return nil + } +} + +type HTTPClient struct { + httpClient *http.Client + userAgent string + apiKey string +} + +var _ Client = (*HTTPClient)(nil) + +func NewHTTPClient(apiKey string, opts ...option) (*HTTPClient, error) { + c := &HTTPClient{ + httpClient: &http.Client{Timeout: defaultHTTPClientTimeout}, + userAgent: defaultUserAgent, + apiKey: apiKey, + } + + for _, opt := range opts { + if err := opt(c); err != nil { + return nil, err + } + } + + return c, nil +} + +type Task struct { + Object string `json:"object"` + ID string `json:"id"` + Created int64 `json:"created"` + TaskType string `json:"task_type"` + Status string `json:"status"` + PromptID string `json:"prompt_id"` + Prompt Prompt `json:"prompt"` + Generations Generations `json:"generations"` +} +type Generations struct { + Data []GenerationData `json:"data"` + Object string `json:"object"` +} +type GenerationData struct { + Created int64 `json:"created"` + Generation Generation `json:"generation"` + GenerationType string `json:"generation_type"` + ID string `json:"id"` +} +type Generation struct { + ImagePath string `json:"image_path"` +} + +type Prompt struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + PromptType string `json:"prompt_type"` + Prompt struct { + Caption string `json:"caption"` + } `json:"prompt"` + ParentGenerationID string `json:"parent_generation_id"` +} + +type GenerateRequest struct { + Prompt GenerateRequestPrompt `json:"prompt"` + TaskType string `json:"task_type"` +} +type GenerateRequestPrompt struct { + BatchSize int32 `json:"batch_size"` + Caption string `json:"caption"` +} + +func (c *HTTPClient) Generate(ctx context.Context, caption string) (*Task, error) { + task := &Task{} + req := &GenerateRequest{ + Prompt: GenerateRequestPrompt{ + BatchSize: defaultBatchSize, + Caption: caption, + }, + TaskType: TaskTypeText2Im, + } + return task, c.request(ctx, "POST", "/tasks", nil, req, task) +} + +type ListTasksResponse struct { + Object string `json:"object"` + Data []Task `json:"data"` +} + +type ListTasksRequest struct { + Limit int32 `json:"limit"` +} + +func (c *HTTPClient) ListTasks(ctx context.Context, req *ListTasksRequest) (*ListTasksResponse, error) { + res := &ListTasksResponse{} + url := "/tasks" + if req != nil { + if req.Limit != 0 { + url += fmt.Sprintf("?limit=%d", req.Limit) + } + } + + return res, c.request(ctx, "GET", url, nil, nil, res) +} + +func (c *HTTPClient) GetTask(ctx context.Context, taskID string) (*Task, error) { + task := &Task{} + return task, c.request(ctx, "GET", "/tasks/"+taskID, nil, nil, task) +} + +func (c *HTTPClient) Download(ctx context.Context, generationID string) (io.ReadCloser, error) { + req, err := c.createRequest(ctx, "/generations/"+generationID+"/download", "GET", nil, nil) + if err != nil { + return nil, fmt.Errorf("creating request: %w", err) + } + + resp, err := c.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("performing request: %w", err) + } + + return resp.Body, nil +} + +// Share makes the generation public and returns the public url +func (c *HTTPClient) Share(ctx context.Context, generationID string) (string, error) { + res := &GenerationData{} + + err := c.request(ctx, "POST", "/generations/"+generationID+"/share", nil, nil, res) + if err != nil { + return "", err + } + + return res.Generation.ImagePath, nil +} + +func (c *HTTPClient) createRequest(ctx context.Context, path, method string, values *url.Values, data interface{}) (*http.Request, error) { + url := baseURL + path + + if values != nil { + url += "?" + values.Encode() + } + + var body io.Reader + if data != nil { + b, err := json.Marshal(data) + if err != nil { + return nil, fmt.Errorf("parsing request data: %w", err) + } + body = bytes.NewReader(b) + } + + req, err := http.NewRequestWithContext(ctx, method, url, body) + if err != nil { + return nil, fmt.Errorf("building request: %w", err) + } + + req.Header.Add("Authorization", "Bearer "+c.apiKey) + req.Header.Add("Content-Type", "application/json") + req.Header.Set("User-Agent", c.userAgent) + + req.Header.Set("Authority", "labs.openai.com") + req.Header.Set("Accept", "*/*") + req.Header.Set("Accept-Language", "en-US,en;q=0.9,de;q=0.8") + req.Header.Set("Cache-Control", "no-cache") + req.Header.Set("Content-Length", "0") + req.Header.Set("Cookie", "") + req.Header.Set("Dnt", "1") + req.Header.Set("Origin", "https://labs.openai.com") + req.Header.Set("Pragma", "no-cache") + req.Header.Set("Sec-Ch-Ua", "\"Chromium\";v=\"104\", \" Not A;Brand\";v=\"99\", \"Google Chrome\";v=\"104\"") + req.Header.Set("Sec-Ch-Ua-Mobile", "?0") + req.Header.Set("Sec-Ch-Ua-Platform", "\"macOS\"") + req.Header.Set("Sec-Fetch-Dest", "empty") + req.Header.Set("Sec-Fetch-Mode", "cors") + req.Header.Set("Sec-Fetch-Site", "same-origin") + req.Header.Set("User-Agent", "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/104.0.0.0 Safari/537.36") + + return req, nil +} + +func (c *HTTPClient) request(ctx context.Context, method, path string, values *url.Values, body interface{}, result interface{}) error { + req, err := c.createRequest(ctx, path, method, values, body) + if err != nil { + return fmt.Errorf("creating request: %w", err) + } + + resp, err := c.httpClient.Do(req) + if err != nil { + return fmt.Errorf("performing request: %w", err) + } + defer resp.Body.Close() + + respBody, _ := io.ReadAll(resp.Body) + + // TODO: improve error handling... + if resp.StatusCode != http.StatusOK { + return Error{ + Message: "unexpected non 200 status code", + StatusCode: resp.StatusCode, + Details: string(respBody), + } + } + + if err = json.Unmarshal(respBody, result); err != nil { + return Error{ + Message: err.Error(), + StatusCode: resp.StatusCode, + Details: string(respBody), + } + } + + return nil +} + +type Error struct { + Message string + StatusCode int + Details string +} + +func (e Error) Error() string { + return fmt.Sprintf("dalle: %s (status: %d, details: %s)", e.Message, e.StatusCode, e.Details) +} diff --git a/internal/telegram/generate.go b/internal/telegram/generate.go index 9fd3336..f292755 100644 --- a/internal/telegram/generate.go +++ b/internal/telegram/generate.go @@ -4,12 +4,13 @@ import ( "bytes" "context" "fmt" + "github.com/LightningTipBot/LightningTipBot/internal/dalle" "github.com/LightningTipBot/LightningTipBot/internal/runtime" "github.com/LightningTipBot/LightningTipBot/internal/telegram/intercept" - "github.com/dillonstreator/dalle" log "github.com/sirupsen/logrus" "github.com/skip2/go-qrcode" tb "gopkg.in/lightningtipbot/telebot.v3" + "io" "os" "time" ) @@ -24,7 +25,7 @@ func (bot *TipBot) generateImages(ctx intercept.Context) (intercept.Context, err return ctx, err } invoice, err := bot.createInvoiceWithEvent(ctx, me, 1, fmt.Sprintf("DALLE2 %s", GetUserStr(user.Telegram)), InvoiceCallbackGenerateDalle, "") - + invoice.Payer = user if err != nil { return ctx, err } @@ -44,20 +45,22 @@ func (bot *TipBot) generateImages(ctx intercept.Context) (intercept.Context, err func (bot *TipBot) generateDalleImages(event Event) { invoiceEvent := event.(*InvoiceEvent) - user := invoiceEvent.User + user := invoiceEvent.Payer if user.Wallet == nil { return } // create the client with the bearer token api key + dalleClient, err := dalle.NewHTTPClient("") // handle err if err != nil { return } + ctx, cancel := context.WithTimeout(context.Background(), time.Minute*5) defer cancel() // generate a task to create an image with a prompt - task, err := dalleClient.Generate(ctx, "monkey printing bitcoin with money printing machine, cyberpunk") + task, err := dalleClient.Generate(ctx, "dogs fighting for bitcoin on sunny island, van gogh style") // handle err // poll the task.ID until status is succeeded @@ -67,7 +70,6 @@ func (bot *TipBot) generateDalleImages(event Event) { t, err = dalleClient.GetTask(ctx, task.ID) // handle err - t.Status = dalle.StatusSucceeded if t.Status == dalle.StatusSucceeded { fmt.Println("task succeeded") @@ -78,19 +80,31 @@ func (bot *TipBot) generateDalleImages(event Event) { fmt.Println("task still pending") } - /* - // download the first generated image - for _, data := range t.Generations.Data { - reader, err := dalleClient.Download(ctx, data.ID) - if err != nil { - continue - } - bot.trySendMessage(user.Telegram, &tb.Photo{File: tb.File{FileReader: reader}, Caption: fmt.Sprintf("Result")}) - }*/ - reader, err := os.OpenFile("image", 0, os.ModePerm) - if err != nil { - panic(err) + + // download the first generated image + for _, data := range t.Generations.Data { + + reader, err := dalleClient.Download(ctx, data.ID) + if err != nil { + return + } + defer reader.Close() + + file, err := os.Create("images/" + data.ID + ".png") + if err != nil { + return + } + defer file.Close() + _, err = io.Copy(file, reader) + if err != nil { + return + } + f, err := os.OpenFile("images/"+data.ID+".png", 0, os.ModePerm) + if err != nil { + return + } + bot.trySendMessage(invoiceEvent.Payer.Telegram, &tb.Photo{File: tb.File{FileReader: f}, Caption: fmt.Sprintf("Result")}) } - bot.trySendMessage(user.Telegram, &tb.Photo{File: tb.File{FileReader: reader}, Caption: fmt.Sprintf("Result")}) + // handle err and close readCloser }