Merge branch 'release/2.4.0' of ssh://git.yunion.io/~qiujian/onecloud into feature/qj-esxi-support-complete

This commit is contained in:
Qiu Jian
2018-11-21 23:35:47 +08:00
76 changed files with 3517 additions and 464 deletions
Generated
+4 -1
View File
@@ -120,7 +120,7 @@
revision = "e59b73d3c2bf1c328ccb78e683c0462fa1a473c7"
[[projects]]
digest = "1:55bf2a4da68caa693d660683e53bab3e651a940cd46354b997ab68bec7920e23"
digest = "1:ae41c49d4812dd31a848045447637fd0f10cac09af75e1d20825a241c626f3fa"
name = "github.com/aws/aws-sdk-go"
packages = [
"aws",
@@ -158,6 +158,8 @@
"service/ec2",
"service/iam",
"service/s3",
"service/s3/s3iface",
"service/s3/s3manager",
"service/sts",
]
pruneopts = "UT"
@@ -1406,6 +1408,7 @@
"github.com/aws/aws-sdk-go/service/ec2",
"github.com/aws/aws-sdk-go/service/iam",
"github.com/aws/aws-sdk-go/service/s3",
"github.com/aws/aws-sdk-go/service/s3/s3manager",
"github.com/bitly/go-simplejson",
"github.com/c-bata/go-prompt",
"github.com/coredns/coredns/core/dnsserver",
+16
View File
@@ -0,0 +1,16 @@
package shell
import (
"yunion.io/x/onecloud/pkg/mcclient"
)
func init() {
type CloudmetaOptions struct {
PROVIDER_ID string `help:"provider_id"`
REGION_ID string `help:"region_id"`
ZONE_ID string `help:"zone_id"`
}
R(&CloudmetaOptions{}, "instance-type-list", "query backend service for its version", func(s *mcclient.ClientSession, args *CloudmetaOptions) error {
return nil
})
}
-32
View File
@@ -31,38 +31,6 @@ func init() {
return nil
})
type CloudproviderCreateOptions struct {
NAME string `help:"Name of cloud provider"`
ACCOUNT string `help:"Account to access the cloud provider, tenantId/subscriptionId for Azure"`
SECRET string `help:"Secret to access the cloud provider, clientId/clientScret for Azure"`
PROVIDER string `help:"Driver for cloud provider" choices:"VMware|Aliyun|Azure"`
AccessURL string `helo:"hello" metavar:"Azure choices: <AzureGermanCloud、AzureChinaCloud、AzureUSGovernmentCloud、AzurePublicCloud>"`
Desc string `help:"Description"`
Enabled bool `help:"Enabled the provider automatically"`
}
R(&CloudproviderCreateOptions{}, "cloud-provider-create", "Create a cloud provider", func(s *mcclient.ClientSession, args *CloudproviderCreateOptions) error {
params := jsonutils.NewDict()
params.Add(jsonutils.NewString(args.NAME), "name")
params.Add(jsonutils.NewString(args.ACCOUNT), "account")
params.Add(jsonutils.NewString(args.SECRET), "secret")
params.Add(jsonutils.NewString(args.PROVIDER), "provider")
if args.Enabled {
params.Add(jsonutils.JSONTrue, "enabled")
}
if len(args.AccessURL) > 0 {
params.Add(jsonutils.NewString(args.AccessURL), "access_url")
}
if len(args.Desc) > 0 {
params.Add(jsonutils.NewString(args.Desc), "description")
}
result, err := modules.Cloudproviders.Create(s, params)
if err != nil {
return err
}
printObject(result)
return nil
})
type CloudproviderUpdateOptions struct {
ID string `help:"ID or Name of cloud provider"`
Name string `help:"New name to update"`
+13
View File
@@ -203,4 +203,17 @@ func initCluster() {
printObject(ret)
return nil
})
R(&o.ClusterRestartAgentsOptions{}, cmdN("restart-agent"), "Restart node agents in cluster", func(s *mcclient.ClientSession, args *o.ClusterRestartAgentsOptions) error {
params, err := args.Params()
if err != nil {
return err
}
ret, err := k8s.Clusters.PerformAction(s, args.ID, "restart-agent", params)
if err != nil {
return err
}
printObject(ret)
return nil
})
}
+12
View File
@@ -111,6 +111,18 @@ func init() {
return nil
})
type SchedpoliciesShowOptions struct {
ID string `help:"ID or name of the sched policy"`
}
R(&SchedpoliciesShowOptions{}, "sched-policy-show", "show details of a sched policy", func(s *mcclient.ClientSession, args *SchedpoliciesShowOptions) error {
result, err := modules.Schedpolicies.Get(s, args.ID, nil)
if err != nil {
return err
}
printObject(result)
return nil
})
type SchedpoliciesEvaluateOptions struct {
ID string `help:"ID or name of the sched policy"`
SERVER string `help:"ID or name of the server"`
+34
View File
@@ -0,0 +1,34 @@
package shell
import (
"yunion.io/x/jsonutils"
"yunion.io/x/onecloud/pkg/mcclient"
"yunion.io/x/onecloud/pkg/mcclient/modules"
"yunion.io/x/onecloud/pkg/mcclient/options"
)
func init() {
type UnderutilizedInstancesListOptions struct {
options.BaseListOptions
}
R(&UnderutilizedInstancesListOptions{}, "underutilized-instances-list", "List underutilized instances", func(s *mcclient.ClientSession, args *UnderutilizedInstancesListOptions) error {
var params *jsonutils.JSONDict
{
var err error
params, err = args.BaseListOptions.Params()
if err != nil {
return err
}
}
result, err := modules.UnderutilizedInstances.List(s, params)
if err != nil {
return err
}
printList(result, modules.UnderutilizedInstances.GetColumns(s))
return nil
})
}
+50 -28
View File
@@ -57,12 +57,16 @@ func (dispatcher *DBModelDispatcher) ContextKeywordPlural() []string {
}
func (dispatcher *DBModelDispatcher) Filter(f appsrv.FilterHandler) appsrv.FilterHandler {
return auth.Authenticate(f)
if consts.IsRbacEnabled() {
return auth.AuthenticateWithDelayDecision(f, true)
} else {
return auth.Authenticate(f)
}
}
func fetchUserCredential(ctx context.Context) mcclient.TokenCredential {
token := auth.FetchUserCredential(ctx)
if token == nil {
if token == nil && !consts.IsRbacEnabled() {
log.Fatalf("user token credential not found?")
}
return token
@@ -479,14 +483,14 @@ func listItems(manager IModelManager, ctx context.Context, userCred mcclient.Tok
if err != nil {
return nil, httperrors.NewGeneralError(err)
}
retConut := len(retList)
retCount := len(retList)
// apply customizeFilters
retList, err = customizeFilters.DoApply(retList)
if err != nil {
return nil, httperrors.NewGeneralError(err)
}
if len(retList) != retConut {
if len(retList) != retCount {
totalCnt = int64(len(retList))
}
paginate := false
@@ -500,7 +504,7 @@ func listItems(manager IModelManager, ctx context.Context, userCred mcclient.Tok
func calculateListResult(data []jsonutils.JSONObject, total, limit, offset int64, paginate bool) *modules.ListResult {
if paginate {
// do offset first
if offset != 0 {
if offset > 0 {
if total > offset {
data = data[offset:]
} else {
@@ -508,11 +512,13 @@ func calculateListResult(data []jsonutils.JSONObject, total, limit, offset int64
}
}
// do limit
if total > limit {
if limit > 0 && total > limit {
data = data[:limit]
}
}
retResult := modules.ListResult{Data: data, Total: int(total), Limit: int(limit), Offset: int(offset)}
return &retResult
}
@@ -725,7 +731,7 @@ func fetchOwnerProjectId(ctx context.Context, manager IModelManager, userCred mc
if consts.IsRbacEnabled() {
result := policy.PolicyManager.Allow(true, userCred,
consts.GetServiceType(), policy.PolicyDelegation, "")
if result == rbacutils.Allow {
if result == rbacutils.AdminAllow {
isAllow = true
}
} else {
@@ -957,6 +963,34 @@ func (dispatcher *DBModelDispatcher) BatchCreate(ctx context.Context, query json
return results, nil
}
func managerPerformCheckCreateData(
manager IModelManager,
ctx context.Context,
userCred mcclient.TokenCredential,
action string,
ownerProjId string,
query jsonutils.JSONObject,
data jsonutils.JSONObject,
) (jsonutils.JSONObject, error) {
body, err := data.(*jsonutils.JSONDict).Get(manager.Keyword())
if err != nil {
return nil, httperrors.NewGeneralError(err)
}
bodyDict := body.(*jsonutils.JSONDict)
var isAllow bool
if consts.IsRbacEnabled() {
isAllow = isClassActionRbacAllowed(manager, userCred, ownerProjId, policy.PolicyActionPerform, action)
} else {
isAllow = manager.AllowPerformCheckCreateData(ctx, userCred, query, data)
}
if !isAllow {
return nil, httperrors.NewForbiddenError("not allow to perform %s", action)
}
return manager.ValidateCreateData(ctx, userCred, ownerProjId, query, bodyDict)
}
func (dispatcher *DBModelDispatcher) PerformClassAction(ctx context.Context, action string, query jsonutils.JSONObject, data jsonutils.JSONObject) (jsonutils.JSONObject, error) {
userCred := fetchUserCredential(ctx)
@@ -969,25 +1003,8 @@ func (dispatcher *DBModelDispatcher) PerformClassAction(ctx context.Context, act
defer lockman.ReleaseClass(ctx, dispatcher.modelManager, ownerProjId)
if action == "check-create-data" {
manager := dispatcher.modelManager
body, err := data.(*jsonutils.JSONDict).Get(manager.Keyword())
if err != nil {
return nil, httperrors.NewGeneralError(err)
}
data := body.(*jsonutils.JSONDict)
var isAllow bool
if consts.IsRbacEnabled() {
isAllow = isClassActionRbacAllowed(manager, userCred, ownerProjId, policy.PolicyActionPerform, action)
} else {
isAllow = manager.AllowPerformCheckCreateData(ctx, userCred, query, data)
}
if !isAllow {
return nil, httperrors.NewForbiddenError("not allow to perform %s", action)
}
return manager.ValidateCreateData(ctx, userCred, ownerProjId, query, data)
return managerPerformCheckCreateData(dispatcher.modelManager,
ctx, userCred, action, ownerProjId, query, data)
}
managerValue := reflect.ValueOf(dispatcher.modelManager)
@@ -1022,8 +1039,8 @@ func objectPerformAction(dispatcher *DBModelDispatcher, model IModel, modelValue
isGeneral := false
funcName := fmt.Sprintf("Perform%s", utils.Kebab2Camel(action, "-"))
funcValue := modelValue.MethodByName(funcName)
if !funcValue.IsValid() || funcValue.IsNil() {
funcValue = modelValue.MethodByName(generalFuncName)
if !funcValue.IsValid() || funcValue.IsNil() {
@@ -1057,7 +1074,12 @@ func objectPerformAction(dispatcher *DBModelDispatcher, model IModel, modelValue
var isAllow bool
if consts.IsRbacEnabled() {
isAllow = isObjectRbacAllowed(dispatcher.modelManager, model, userCred, policy.PolicyActionPerform, action)
if model == nil {
ownerProjId, _ := fetchOwnerProjectId(ctx, dispatcher.modelManager, userCred, data)
isAllow = isClassActionRbacAllowed(dispatcher.modelManager, userCred, ownerProjId, policy.PolicyActionPerform, action)
} else {
isAllow = isObjectRbacAllowed(dispatcher.modelManager, model, userCred, policy.PolicyActionPerform, action)
}
} else {
allowFuncName := "Allow" + funcName
allowFuncValue := modelValue.MethodByName(allowFuncName)
+5 -5
View File
@@ -91,7 +91,7 @@ func getQuotaHanlder(ctx context.Context, w http.ResponseWriter, r *http.Request
if consts.IsRbacEnabled() {
result := policy.PolicyManager.Allow(true, userCred, consts.GetServiceType(),
policy.PolicyDelegation, policy.PolicyActionGet)
isAllow = result == rbacutils.Allow
isAllow = result == rbacutils.AdminAllow
} else {
isAllow = userCred.IsSystemAdmin()
}
@@ -101,7 +101,7 @@ func getQuotaHanlder(ctx context.Context, w http.ResponseWriter, r *http.Request
}
if consts.IsRbacEnabled() {
if policy.PolicyManager.Allow(true, userCred, consts.GetServiceType(),
"quotas", policy.PolicyActionGet) != rbacutils.Allow {
"quotas", policy.PolicyActionGet) != rbacutils.AdminAllow {
httperrors.ForbiddenError(w, "not allow to query quota")
return
}
@@ -138,7 +138,7 @@ func setQuotaHanlder(ctx context.Context, w http.ResponseWriter, r *http.Request
var isAllow bool
if consts.IsRbacEnabled() {
isAllow = policy.PolicyManager.Allow(true, userCred, consts.GetServiceType(),
"quotas", policy.PolicyActionUpdate) == rbacutils.Allow
"quotas", policy.PolicyActionUpdate) == rbacutils.AdminAllow
} else {
isAllow = userCred.IsSystemAdmin()
}
@@ -201,7 +201,7 @@ func checkQuotaHanlder(ctx context.Context, w http.ResponseWriter, r *http.Reque
isAllow := false
if consts.IsRbacEnabled() {
isAllow = policy.PolicyManager.Allow(true, userCred, consts.GetServiceType(),
policy.PolicyDelegation, policy.PolicyActionGet) == rbacutils.Allow
policy.PolicyDelegation, policy.PolicyActionGet) == rbacutils.AdminAllow
} else {
isAllow = userCred.IsSystemAdmin()
}
@@ -211,7 +211,7 @@ func checkQuotaHanlder(ctx context.Context, w http.ResponseWriter, r *http.Reque
}
if consts.IsRbacEnabled() {
if policy.PolicyManager.Allow(true, userCred, consts.GetServiceType(),
"quotas", policy.PolicyActionGet) != rbacutils.Allow {
"quotas", policy.PolicyActionGet) != rbacutils.AdminAllow {
httperrors.ForbiddenError(w, "not allow to query quota")
return
}
+33 -12
View File
@@ -28,14 +28,22 @@ func isListRbacAllowedInternal(manager IModelManager, resource string, userCred
}
result := policy.PolicyManager.Allow(false, userCred, consts.GetServiceType(),
resource, policy.PolicyActionList)
log.Debugf("allow list for non-admin %s %v ownerId: %s", result, requireAdmin, ownerId)
if (result == rbacutils.OwnerAllow && !requireAdmin) || (result == rbacutils.Allow && requireAdmin) {
switch {
case result == rbacutils.GuestAllow:
return true
case result == rbacutils.UserAllow && userCred != nil && userCred.IsValid():
return true
case result == rbacutils.OwnerAllow && !requireAdmin:
return true
case result == rbacutils.AdminAllow && requireAdmin:
return true
}
result = policy.PolicyManager.Allow(true, userCred, consts.GetServiceType(),
resource, policy.PolicyActionList)
log.Debugf("allow list for admin %s %v ownerId: %s", result, requireAdmin, ownerId)
return result == rbacutils.Allow
return result == rbacutils.AdminAllow
}
func isJointListRbacAllowed(manager IJointModelManager, userCred mcclient.TokenCredential, isAdminMode bool) bool {
@@ -54,16 +62,23 @@ func isClassActionRbacAllowed(manager IModelManager, userCred mcclient.TokenCred
} else {
requireAdmin = true
}
// if !requireAdmin {
result := policy.PolicyManager.Allow(false, userCred, consts.GetServiceType(),
manager.KeywordPlural(), action, extra...)
if result == rbacutils.Allow || (!requireAdmin && result == rbacutils.OwnerAllow) {
switch {
case result == rbacutils.GuestAllow:
return true
case result == rbacutils.UserAllow && userCred != nil && userCred.IsValid():
return true
case result == rbacutils.OwnerAllow && !requireAdmin:
return true
case result == rbacutils.AdminAllow && requireAdmin:
return true
}
// }
result = policy.PolicyManager.Allow(true, userCred, consts.GetServiceType(),
manager.KeywordPlural(), action, extra...)
return result == rbacutils.Allow
return result == rbacutils.AdminAllow
}
func isObjectRbacAllowed(manager IModelManager, model IModel, userCred mcclient.TokenCredential, action string, extra ...string) bool {
@@ -85,16 +100,22 @@ func isObjectRbacAllowed(manager IModelManager, model IModel, userCred mcclient.
requireAdmin = true
}
//if !requireAdmin {
result := policy.PolicyManager.Allow(false, userCred, consts.GetServiceType(),
manager.KeywordPlural(), action, extra...)
if result == rbacutils.Allow || (!requireAdmin && result == rbacutils.OwnerAllow && isOwner) {
switch {
case result == rbacutils.GuestAllow:
return true
case result == rbacutils.UserAllow && userCred != nil && userCred.IsValid():
return true
case result == rbacutils.OwnerAllow && isOwner && !requireAdmin:
return true
case result == rbacutils.AdminAllow && (requireAdmin || isOwner):
return true
}
//}
result = policy.PolicyManager.Allow(true, userCred, consts.GetServiceType(),
manager.KeywordPlural(), action, extra...)
return result == rbacutils.Allow
return result == rbacutils.AdminAllow
}
func isJointObjectRbacAllowed(manager IJointModelManager, item IJointModel, userCred mcclient.TokenCredential, action string, extra ...string) bool {
+14 -5
View File
@@ -89,10 +89,6 @@ func (manager *STaskManager) FilterByName(q *sqlchemy.SQuery, name string) *sqlc
return q
}
func (manager *STaskManager) FilterByOwner(q *sqlchemy.SQuery, owner string) *sqlchemy.SQuery {
return q
}
func (manager *STaskManager) AllowPerformAction(ctx context.Context, userCred mcclient.TokenCredential, action string, query jsonutils.JSONObject, data jsonutils.JSONObject) bool {
return true
}
@@ -105,6 +101,19 @@ func (manager *STaskManager) PerformAction(ctx context.Context, userCred mcclien
return resp, nil
}
func (manager *STaskManager) GetOwnerId(userCred mcclient.IIdentityProvider) string {
return userCred.GetProjectId()
}
func (self *STask) GetOwnerProjectId() string {
return self.UserCred.GetProjectId()
}
func (manager *STaskManager) FilterByOwner(q *sqlchemy.SQuery, owner string) *sqlchemy.SQuery {
q = q.Contains("user_cred", owner)
return q
}
func (self *STask) AllowGetDetails(ctx context.Context, userCred mcclient.TokenCredential, query jsonutils.JSONObject) bool {
return userCred.IsSystemAdmin() || userCred.GetProjectId() == self.UserCred.GetProjectId()
}
@@ -414,7 +423,7 @@ func execITask(taskValue reflect.Value, task *STask, odata jsonutils.JSONObject,
return
}
log.Debugf("Call %s %s: %s with %s", task.TaskName, stageName, funcValue, params)
log.Debugf("Call %s %s", task.TaskName, stageName)
funcValue.Call(params)
// call save request context
+1 -1
View File
@@ -40,7 +40,7 @@ type Options struct {
SslCertfile string `help:"ssl certification file"`
SslKeyfile string `help:"ssl certification key file"`
EnableRbac bool `help:"Switch on Role-based Access Control" default:"false"`
EnableRbac bool `help:"Switch on Role-based Access Control" default:"true"`
RbacDebug bool `help:"turn on rbac debug log" default:"false"`
RbacPolicySyncPeriodSeconds int `help:"policy sync interval in seconds, default 15 minutes" default:"900"`
RbacPolicySyncFailedRetrySeconds int `help:"seconds to wait after a failed sync, default 30 seconds" default:"30"`
+62 -10
View File
@@ -30,8 +30,37 @@ const (
var (
PolicyManager *SPolicyManager
PolicyFailedRetryInterval = 15 * time.Second
PolicyRefreshInterval = 15 * time.Minute
defaultRules = []rbacutils.SRbacRule{
{
Resource: "tasks",
Action: PolicyActionPerform,
Result: rbacutils.UserAllow,
},
{
Service: "compute",
Resource: "zones",
Action: PolicyActionList,
Result: rbacutils.UserAllow,
},
{
Service: "compute",
Resource: "zones",
Action: PolicyActionGet,
Result: rbacutils.UserAllow,
},
{
Service: "compute",
Resource: "cloudregions",
Action: PolicyActionList,
Result: rbacutils.UserAllow,
},
{
Service: "compute",
Resource: "cloudregions",
Action: PolicyActionGet,
Result: rbacutils.UserAllow,
},
}
)
func init() {
@@ -41,7 +70,11 @@ func init() {
type SPolicyManager struct {
policies map[string]rbacutils.SRbacPolicy
adminPolicies map[string]rbacutils.SRbacPolicy
defaultPolicy *rbacutils.SRbacPolicy
lastSync time.Time
failedRetryInterval time.Duration
refreshInterval time.Duration
}
func parseJsonPolicy(obj jsonutils.JSONObject) (string, rbacutils.SRbacPolicy, error) {
@@ -114,8 +147,13 @@ func fetchPolicies() (map[string]rbacutils.SRbacPolicy, map[string]rbacutils.SRb
func (manager *SPolicyManager) start(refreshInterval time.Duration, retryInterval time.Duration) {
log.Infof("PolicyManager start to fetch policies ...")
PolicyRefreshInterval = refreshInterval
PolicyFailedRetryInterval = retryInterval
manager.refreshInterval = refreshInterval
manager.failedRetryInterval = retryInterval
if len(defaultRules) > 0 {
manager.defaultPolicy = &rbacutils.SRbacPolicy{
Rules: rbacutils.CompactRules(defaultRules),
}
}
manager.sync()
}
@@ -124,13 +162,14 @@ func (manager *SPolicyManager) sync() {
policies, adminPolicies, err := fetchPolicies()
if err != nil {
log.Errorf("sync policy fail %s", err)
time.AfterFunc(PolicyFailedRetryInterval, manager.sync)
time.AfterFunc(manager.failedRetryInterval, manager.sync)
return
}
manager.policies = policies
manager.adminPolicies = adminPolicies
manager.lastSync = time.Now()
time.AfterFunc(PolicyRefreshInterval, manager.sync)
time.AfterFunc(manager.refreshInterval, manager.sync)
}
func (manager *SPolicyManager) Allow(isAdmin bool, userCred mcclient.TokenCredential, service string, resource string, action string, extra ...string) rbacutils.TRbacResult {
@@ -148,7 +187,13 @@ func (manager *SPolicyManager) Allow(isAdmin bool, userCred mcclient.TokenCreden
currentPriv := rbacutils.Deny
for _, p := range policies {
result := p.Allow(userCredJson, service, resource, action, extra...)
if result.IsHigherPrivilege(currentPriv) {
if currentPriv.StricterThan(result) {
currentPriv = result
}
}
if manager.defaultPolicy != nil {
result := manager.defaultPolicy.Allow(userCredJson, service, resource, action, extra...)
if currentPriv.StricterThan(result) {
currentPriv = result
}
}
@@ -165,8 +210,10 @@ func (manager *SPolicyManager) explainPolicy(userCred mcclient.TokenCredential,
}
isAdmin, _ := policySeq[0].Bool()
if !consts.IsRbacEnabled() {
if !isAdmin || (isAdmin && userCred.IsSystemAdmin()) {
return rbacutils.Allow, nil
if !isAdmin {
return rbacutils.OwnerAllow, nil
} else if isAdmin && userCred.IsSystemAdmin() {
return rbacutils.AdminAllow, nil
} else {
return rbacutils.Deny, httperrors.NewForbiddenError("operation not allowed")
}
@@ -186,7 +233,8 @@ func (manager *SPolicyManager) explainPolicy(userCred mcclient.TokenCredential,
}
if len(policySeq) > 4 {
for i := 4; i < len(policySeq); i += 1 {
extra[i-4], _ = policySeq[i].GetString()
ev, _ := policySeq[i].GetString()
extra = append(extra, ev)
}
}
@@ -210,6 +258,10 @@ func (manager *SPolicyManager) ExplainRpc(userCred mcclient.TokenCredential, par
}
func (manager *SPolicyManager) IsAdminCapable(userCred mcclient.TokenCredential) bool {
if !consts.IsRbacEnabled() && userCred.IsSystemAdmin() {
return true
}
userCredJson := userCred.ToJson()
for _, p := range manager.adminPolicies {
match, _ := conditionparser.Eval(p.Condition, userCredJson)
+34 -26
View File
@@ -5,6 +5,7 @@ import (
"fmt"
"yunion.io/x/onecloud/pkg/httperrors"
"yunion.io/x/onecloud/pkg/util/httputils"
)
var returnHttpError = true
@@ -54,65 +55,72 @@ func (ve *ValidateError) Error() string {
// TODO let each validator provide the error
func newMissingKeyError(key string) error {
msg := fmt.Sprintf("missing %q", key)
return newError(ERR_MISSING_KEY, msg)
return newError(ERR_MISSING_KEY, "missing %q", key)
}
func newGeneralError(key string, err error) error {
msg := fmt.Sprintf("general error for %q: %s", key, err)
return newError(ERR_GENERAL, msg)
return newError(ERR_GENERAL, "general error for %q: %s", key, err)
}
func newInvalidTypeError(key string, typ string, err error) error {
msg := fmt.Sprintf("expecting %s type for %q: %s", typ, key, err)
return newError(ERR_INVALID_TYPE, msg)
return newError(ERR_INVALID_TYPE, "expecting %s type for %q: %s", typ, key, err)
}
func newInvalidChoiceError(key string, choices Choices, choice string) error {
msg := fmt.Sprintf("invalid %q, want %s, got %s", key, choices, choice)
return newError(ERR_INVALID_CHOICE, msg)
return newError(ERR_INVALID_CHOICE, "invalid %q, want %s, got %s", key, choices, choice)
}
func newNotInRangeError(key string, value, lower, upper int64) error {
msg := fmt.Sprintf("invalid %q: %d, want [%d,%d]", key, value, lower, upper)
return newError(ERR_NOT_IN_RANGE, msg)
return newError(ERR_NOT_IN_RANGE, "invalid %q: %d, want [%d,%d]", key, value, lower, upper)
}
func newInvalidValueError(key string, value string) error {
msg := fmt.Sprintf("invalid %q: %s", key, value)
return newError(ERR_INVALID_VALUE, msg)
return newError(ERR_INVALID_VALUE, "invalid %q: %s", key, value)
}
func newInvalidStructError(key string, err error) error {
errFmt := "invalid %q: "
params := []interface{}{key}
jsonClientErr, ok := err.(*httputils.JSONClientError)
if ok {
errFmt += jsonClientErr.Data.Id
for _, f := range jsonClientErr.Data.Fields {
params = append(params, f)
}
}
return newError(ERR_INVALID_VALUE, errFmt, params...)
}
func newModelManagerError(modelKeyword string) error {
msg := fmt.Sprintf("internal error: getting model manager for %q failed",
modelKeyword)
return newError(ERR_MODEL_MANAGER, msg)
return newError(ERR_MODEL_MANAGER, "failed getting model manager for %q", modelKeyword)
}
func newModelNotFoundError(modelKeyword, idOrName string, err error) error {
msg := fmt.Sprintf("cannot find %q with id/name %q",
modelKeyword, idOrName)
errFmt := "cannot find %q with id/name %q"
params := []interface{}{modelKeyword, idOrName}
if err != sql.ErrNoRows {
msg += ": " + err.Error()
errFmt += ": %s"
params = append(params, err.Error())
}
return newError(ERR_MODEL_NOT_FOUND, msg)
return newError(ERR_MODEL_NOT_FOUND, errFmt, params...)
}
func newError(typ ErrType, msg string) error {
err := &ValidateError{
ErrType: typ,
Msg: msg,
}
func newError(typ ErrType, errFmt string, params ...interface{}) error {
errFmt = fmt.Sprintf("%s: %s", typ, errFmt)
if returnHttpError {
switch typ {
case ERR_SUCCESS:
return nil
case ERR_GENERAL, ERR_MODEL_MANAGER:
return httperrors.NewInternalServerError(msg)
return httperrors.NewInternalServerError(errFmt, params...)
default:
return httperrors.NewInputParameterError(msg)
return httperrors.NewInputParameterError(errFmt, params...)
}
}
err := &ValidateError{
ErrType: typ,
Msg: fmt.Sprintf(errFmt, params...),
}
return err
}
+1 -1
View File
@@ -511,7 +511,7 @@ func (v *ValidatorStruct) Validate(data *jsonutils.JSONDict) error {
if valueValidator, ok := v.Value.(IValidatorBase); ok {
err = valueValidator.Validate(data)
if err != nil {
return newInvalidValueError(v.Key, err.Error())
return newInvalidStructError(v.Key, err)
}
}
data.Set(v.Key, jsonutils.Marshal(v.Value))
-2
View File
@@ -257,12 +257,10 @@ type ICloudDisk interface {
type ICloudSnapshot interface {
ICloudResource
GetManagerId() string
GetSize() int32
GetDiskId() string
GetDiskType() string
Delete() error
GetRegionId() string
}
type ICloudVpc interface {
+4 -25
View File
@@ -5,6 +5,8 @@ import (
"fmt"
"time"
"yunion.io/x/onecloud/pkg/util/ansible"
"yunion.io/x/jsonutils"
"yunion.io/x/log"
"yunion.io/x/onecloud/pkg/cloudcommon/db"
@@ -158,7 +160,7 @@ func (self *SAwsGuestDriver) RequestDeployGuestOnHost(ctx context.Context, guest
return nil, err
}
data := fetchIVMinfo(desc, iVM, guest.Id, DEFAULT_USER, passwd, action)
data := fetchIVMinfo(desc, iVM, guest.Id, ansible.PUBLIC_CLOUD_ANSIBLE_USER, passwd, action)
return data, nil
})
case "rebuild":
@@ -210,7 +212,7 @@ func (self *SAwsGuestDriver) RequestDeployGuestOnHost(ctx context.Context, guest
}
}
data := fetchIVMinfo(desc, iVM, guest.Id, DEFAULT_USER, passwd, action)
data := fetchIVMinfo(desc, iVM, guest.Id, ansible.PUBLIC_CLOUD_ANSIBLE_USER, passwd, action)
return data, nil
})
@@ -288,29 +290,6 @@ func (self *SAwsGuestDriver) OnGuestDeployTaskDataReceived(ctx context.Context,
return nil
}
func (self *SAwsGuestDriver) RequestDiskSnapshot(ctx context.Context, guest *models.SGuest, task taskman.ITask, snapshotId, diskId string) error {
iDisk, _ := models.DiskManager.FetchById(diskId)
disk := iDisk.(*models.SDisk)
providerDisk, err := disk.GetIDisk()
if err != nil {
return err
}
iSnapshot, _ := models.SnapshotManager.FetchById(snapshotId)
snapshot := iSnapshot.(*models.SSnapshot)
taskman.LocalTaskRun(task, func() (jsonutils.JSONObject, error) {
cloudSnapshot, err := providerDisk.CreateISnapshot(ctx, snapshot.Name, "")
if err != nil {
return nil, err
}
res := jsonutils.NewDict()
res.Set("snapshot_id", jsonutils.NewString(cloudSnapshot.GetId()))
res.Set("manager_id", jsonutils.NewString(cloudSnapshot.GetManagerId()))
res.Set("cloudregion_id", jsonutils.NewString(cloudSnapshot.GetRegionId()))
return res, nil
})
return nil
}
func init() {
driver := SAwsGuestDriver{}
models.RegisterGuestDriver(&driver)
+4 -7
View File
@@ -9,6 +9,7 @@ import (
"yunion.io/x/onecloud/pkg/cloudprovider"
"yunion.io/x/onecloud/pkg/httperrors"
"yunion.io/x/onecloud/pkg/mcclient"
"yunion.io/x/onecloud/pkg/util/ansible"
"yunion.io/x/onecloud/pkg/util/seclib2"
"yunion.io/x/pkg/utils"
@@ -22,10 +23,6 @@ type SAzureGuestDriver struct {
SManagedVirtualizedGuestDriver
}
const (
DEFAULT_USER = "yunion"
)
func init() {
driver := SAzureGuestDriver{}
models.RegisterGuestDriver(&driver)
@@ -169,7 +166,7 @@ func (self *SAzureGuestDriver) RequestDeployGuestOnHost(ctx context.Context, gue
return nil, err
}
data := fetchIVMinfo(desc, iVM, guest.Id, DEFAULT_USER, passwd, action)
data := fetchIVMinfo(desc, iVM, guest.Id, ansible.PUBLIC_CLOUD_ANSIBLE_USER, passwd, action)
return data, nil
}
})
@@ -192,7 +189,7 @@ func (self *SAzureGuestDriver) RequestDeployGuestOnHost(ctx context.Context, gue
if err != nil {
return nil, err
}
data := fetchIVMinfo(desc, iVM, guest.Id, DEFAULT_USER, passwd, action)
data := fetchIVMinfo(desc, iVM, guest.Id, ansible.PUBLIC_CLOUD_ANSIBLE_USER, passwd, action)
return data, nil
})
} else if action == "rebuild" {
@@ -209,7 +206,7 @@ func (self *SAzureGuestDriver) RequestDeployGuestOnHost(ctx context.Context, gue
}
log.Debugf("VMrebuildRoot %s, and status is ready", iVM.GetGlobalId())
data := fetchIVMinfo(desc, iVM, guest.Id, DEFAULT_USER, passwd, action)
data := fetchIVMinfo(desc, iVM, guest.Id, ansible.PUBLIC_CLOUD_ANSIBLE_USER, passwd, action)
return data, nil
})
@@ -368,8 +368,6 @@ func (self *SManagedVirtualizedGuestDriver) RequestDiskSnapshot(ctx context.Cont
}
res := jsonutils.NewDict()
res.Set("snapshot_id", jsonutils.NewString(cloudSnapshot.GetId()))
res.Set("manager_id", jsonutils.NewString(cloudSnapshot.GetManagerId()))
res.Set("cloudregion_id", jsonutils.NewString(cloudSnapshot.GetRegionId()))
return res, nil
})
return nil
-27
View File
@@ -345,30 +345,3 @@ func (self *SQcloudGuestDriver) OnGuestDeployTaskDataReceived(ctx context.Contex
func (self *SQcloudGuestDriver) AllowReconfigGuest() bool {
return true
}
func (self *SQcloudGuestDriver) RequestDiskSnapshot(ctx context.Context, guest *models.SGuest, task taskman.ITask, snapshotId, diskId string) error {
iDisk, _ := models.DiskManager.FetchById(diskId)
disk := iDisk.(*models.SDisk)
providerDisk, err := disk.GetIDisk()
if err != nil {
return err
}
iSnapshot, _ := models.SnapshotManager.FetchById(snapshotId)
snapshot := iSnapshot.(*models.SSnapshot)
taskman.LocalTaskRun(task, func() (jsonutils.JSONObject, error) {
cloudSnapshot, err := providerDisk.CreateISnapshot(ctx, snapshot.Name, "")
if err != nil {
return nil, err
}
res := jsonutils.NewDict()
res.Set("snapshot_id", jsonutils.NewString(cloudSnapshot.GetId()))
res.Set("manager_id", jsonutils.NewString(cloudSnapshot.GetManagerId()))
cloudRegion, err := models.CloudregionManager.FetchByExternalId("Aliyun/" + cloudSnapshot.GetRegionId())
if err != nil {
return nil, fmt.Errorf("Cloud region not found? %s", err)
}
res.Set("cloudregion_id", jsonutils.NewString(cloudRegion.GetId()))
return res, nil
})
return nil
}
+16
View File
@@ -188,6 +188,22 @@ func (self *SAwsHostDriver) RequestResizeDiskOnHost(ctx context.Context, host *m
return nil
}
func (self *SAwsHostDriver) RequestResetDisk(ctx context.Context, host *models.SHost, disk *models.SDisk, params *jsonutils.JSONDict, task taskman.ITask) error {
iDisk, err := disk.GetIDisk()
if err != nil {
return err
}
snapshotId, err := params.GetString("snapshot_id")
if err != nil {
return err
}
taskman.LocalTaskRun(task, func() (jsonutils.JSONObject, error) {
err := iDisk.Reset(snapshotId)
return nil, err
})
return nil
}
func init() {
driver := SAwsHostDriver{}
models.RegisterHostDriver(&driver)
+3 -1
View File
@@ -147,7 +147,8 @@ func (manager *SCloudaccountManager) ValidateCreateData(ctx context.Context, use
if err == cloudprovider.ErrNoSuchProvder {
return nil, httperrors.NewResourceNotFoundError("no such provider %s", provider)
}
return nil, httperrors.NewInvalidCredentialError("invalid cloud account info")
log.Debugf("ValidateCreateData %s", err.Error())
return nil, httperrors.NewInputParameterError("invalid cloud account info")
}
return manager.SEnabledStatusStandaloneResourceBaseManager.ValidateCreateData(ctx, userCred, ownerProjId, query, data)
@@ -386,6 +387,7 @@ func (self *SCloudaccount) ImportSubAccount(ctx context.Context, userCred mcclie
newCloudprovider.CloudaccountId = self.Id
newCloudprovider.Provider = self.Provider
newCloudprovider.Enabled = true
newCloudprovider.Status = CLOUD_PROVIDER_CONNECTED
newCloudprovider.Name = subAccount.Name
if !autoCreateProject {
newCloudprovider.ProjectId = auth.AdminCredential().GetProjectId()
+11 -3
View File
@@ -130,7 +130,7 @@ func (self *SCloudprovider) ValidateUpdateData(ctx context.Context, userCred mcc
}
func (self *SCloudproviderManager) ValidateCreateData(ctx context.Context, userCred mcclient.TokenCredential, ownerProjId string, query jsonutils.JSONObject, data *jsonutils.JSONDict) (*jsonutils.JSONDict, error) {
return nil, httperrors.NewUnsupportOperationError("Not support create cloudprovider, please considir create cloudaccount")
return nil, httperrors.NewUnsupportOperationError("Directly creating cloudprovider is not supported, create cloudaccount instead")
}
func (self *SCloudprovider) getPassword() (string, error) {
@@ -400,7 +400,15 @@ func (self *SCloudprovider) getAccount() (SAccount, error) {
cloudaccount := self.GetCloudaccount()
if cloudaccount == nil {
return account, fmt.Errorf("fail to find cloudaccount???")
// legacy mode
passwd, err := self.getPassword()
if err != nil {
return account, err
}
account.Account = self.Account
account.AccessUrl = self.AccessUrl
account.Secret = passwd
return account, nil // fmt.Errorf("fail to find cloudaccount???")
}
passwd, err := cloudaccount.getPassword()
@@ -429,7 +437,7 @@ func (self *SCloudprovider) SaveSysInfo(info jsonutils.JSONObject) {
func (manager *SCloudproviderManager) FetchCloudproviderById(providerId string) *SCloudprovider {
providerObj, err := manager.FetchById(providerId)
if err != nil {
log.Errorf("%s", err)
log.Errorf("fetch cloud provider %s: %s", providerId, err)
return nil
}
return providerObj.(*SCloudprovider)
+1 -1
View File
@@ -288,7 +288,7 @@ func (self *SCloudregion) PerformDefaultVpc(ctx context.Context, userCred mcclie
func (manager *SCloudregionManager) FetchRegionById(id string) *SCloudregion {
obj, err := manager.FetchById(id)
if err != nil {
log.Errorf("%s", err)
log.Errorf("region %s %s", id, err)
return nil
}
return obj.(*SCloudregion)
+15 -3
View File
@@ -962,7 +962,7 @@ func (manager *SHostManager) getHostsByZoneProvider(zone *SZone, provider *SClou
return hosts, nil
}
func (manager *SHostManager) SyncHosts(ctx context.Context, userCred mcclient.TokenCredential, provider *SCloudprovider, zone *SZone, hosts []cloudprovider.ICloudHost) ([]SHost, []cloudprovider.ICloudHost, compare.SyncResult) {
func (manager *SHostManager) SyncHosts(ctx context.Context, userCred mcclient.TokenCredential, provider *SCloudprovider, zone *SZone, hosts []cloudprovider.ICloudHost, projectSync bool) ([]SHost, []cloudprovider.ICloudHost, compare.SyncResult) {
localHosts := make([]SHost, 0)
remoteHosts := make([]cloudprovider.ICloudHost, 0)
syncResult := compare.SyncResult{}
@@ -1006,7 +1006,7 @@ func (manager *SHostManager) SyncHosts(ctx context.Context, userCred mcclient.To
}
}
for i := 0; i < len(commondb); i += 1 {
err = commondb[i].syncWithCloudHost(commonext[i])
err = commondb[i].syncWithCloudHost(commonext[i], projectSync)
if err != nil {
syncResult.UpdateError(err)
} else {
@@ -1029,7 +1029,7 @@ func (manager *SHostManager) SyncHosts(ctx context.Context, userCred mcclient.To
return localHosts, remoteHosts, syncResult
}
func (self *SHost) syncWithCloudHost(extHost cloudprovider.ICloudHost) error {
func (self *SHost) syncWithCloudHost(extHost cloudprovider.ICloudHost, projectSync bool) error {
_, err := self.GetModelManager().TableSpec().Update(self, func() error {
self.Name = extHost.GetName()
@@ -1060,6 +1060,13 @@ func (self *SHost) syncWithCloudHost(extHost cloudprovider.ICloudHost) error {
if err != nil {
log.Errorf("syncWithCloudZone error %s", err)
}
if projectSync {
if err := HostManager.ClearSchedDescCache(self.Id); err != nil {
log.Errorf("ClearSchedDescCache for host %s error %v", self.Name, err)
}
}
return err
}
@@ -1110,6 +1117,11 @@ func (manager *SHostManager) newFromCloudHost(extHost cloudprovider.ICloudHost,
log.Errorf("newFromCloudHost fail %s", err)
return nil, err
}
if err := manager.ClearSchedDescCache(host.Id); err != nil {
log.Errorf("ClearSchedDescCache for host %s error %v", host.Name, err)
}
return &host, nil
}
+5 -4
View File
@@ -14,6 +14,7 @@ import (
"yunion.io/x/onecloud/pkg/cloudcommon/db"
"yunion.io/x/onecloud/pkg/cloudcommon/validators"
"yunion.io/x/onecloud/pkg/httperrors"
"yunion.io/x/onecloud/pkg/mcclient"
)
@@ -33,16 +34,16 @@ func (aclEntry *SLoadbalancerAclEntry) Validate(data *jsonutils.JSONDict) error
} else {
ip := net.ParseIP(aclEntry.Cidr).To4()
if ip == nil {
return fmt.Errorf("invalid addr %s", aclEntry.Cidr)
return httperrors.NewInputParameterError("invalid addr %s", aclEntry.Cidr)
}
}
if commentLimit := 128; len(aclEntry.Comment) > commentLimit {
return fmt.Errorf("comment too long (%d>=%d)",
return httperrors.NewInputParameterError("comment too long (%d>=%d)",
len(aclEntry.Comment), commentLimit)
}
for _, r := range aclEntry.Comment {
if !unicode.IsPrint(r) {
return fmt.Errorf("comment contains non-printable char: %v", r)
return httperrors.NewInputParameterError("comment contains non-printable char: %v", r)
}
}
return nil
@@ -68,7 +69,7 @@ func (aclEntries *SLoadbalancerAclEntries) Validate(data *jsonutils.JSONDict) er
}
if _, ok := found[aclEntry.Cidr]; ok {
// error so that the user has a chance to deal with comments
return fmt.Errorf("acl cidr duplicate %s", aclEntry.Cidr)
return httperrors.NewInputParameterError("acl cidr duplicate %s", aclEntry.Cidr)
}
found[aclEntry.Cidr] = true
}
+1 -1
View File
@@ -154,7 +154,7 @@ func (p *SLoadbalancerAgentParamsTelegraf) Validate(data *jsonutils.JSONDict) er
if p.InfluxDbOutputUrl != "" {
_, err := url.Parse(p.InfluxDbOutputUrl)
if err != nil {
return err
return httperrors.NewInputParameterError("telegraf params: invalid influxdb url: %s", err)
}
}
if p.HaproxyInputInterval <= 0 {
@@ -43,6 +43,7 @@ type SLoadbalancerListenerRule struct {
func loadbalancerListenerRuleCheckUniqueness(ctx context.Context, lbls *SLoadbalancerListener, domain, path string) error {
q := LoadbalancerListenerRuleManager.Query().
IsFalse("pending_deleted").
Equals("listener_id", lbls.Id).
Equals("domain", domain).
Equals("path", path)
+24 -10
View File
@@ -104,6 +104,7 @@ type SLoadbalancerListener struct {
func (man *SLoadbalancerListenerManager) checkListenerUniqueness(ctx context.Context, lb *SLoadbalancer, listenerType string, listenerPort int64) error {
q := man.Query().
IsFalse("pending_deleted").
Equals("loadbalancer_id", lb.Id).
Equals("listener_port", listenerPort)
switch listenerType {
@@ -260,16 +261,8 @@ func (man *SLoadbalancerListenerManager) ValidateCreateData(ctx context.Context,
}
}
}
{
if aclStatusV.Value == LB_BOOL_ON {
acl := aclV.Model.(*SLoadbalancerAcl)
if acl == nil {
return nil, fmt.Errorf("missing acl")
}
if len(aclTypeV.Value) == 0 {
return nil, fmt.Errorf("missing acl_type")
}
}
if err := man.validateAcl(aclStatusV, aclTypeV, aclV, data); err != nil {
return nil, err
}
return man.SVirtualResourceBaseManager.ValidateCreateData(ctx, userCred, ownerProjId, query, data)
}
@@ -287,6 +280,20 @@ func (man *SLoadbalancerListenerManager) checkTypeV(listenerType string) validat
return nil
}
func (man *SLoadbalancerListenerManager) validateAcl(aclStatusV *validators.ValidatorStringChoices, aclTypeV *validators.ValidatorStringChoices, aclV *validators.ValidatorModelIdOrName, data *jsonutils.JSONDict) error {
if aclStatusV.Value == LB_BOOL_ON {
if aclV.Model == nil {
return httperrors.NewInputParameterError("missing acl")
}
if len(aclTypeV.Value) == 0 {
return httperrors.NewInputParameterError("missing acl_type")
}
} else {
data.Set("acl_id", jsonutils.NewString(""))
}
return nil
}
func (lblis *SLoadbalancerListener) AllowPerformStatus(ctx context.Context, userCred mcclient.TokenCredential, query jsonutils.JSONObject, data jsonutils.JSONObject) bool {
return lblis.IsOwner(userCred) || userCred.IsSystemAdmin()
}
@@ -295,7 +302,11 @@ func (lblis *SLoadbalancerListener) ValidateUpdateData(ctx context.Context, user
ownerProjId := lblis.GetOwnerProjectId()
backendGroupV := validators.NewModelIdOrNameValidator("backend_group", "loadbalancerbackendgroup", ownerProjId)
aclStatusV := validators.NewStringChoicesValidator("acl_status", LB_BOOL_VALUES)
aclStatusV.Default(lblis.AclStatus)
aclTypeV := validators.NewStringChoicesValidator("acl_type", LB_ACL_TYPES)
if LB_ACL_TYPES.Has(lblis.AclType) {
aclTypeV.Default(lblis.AclType)
}
aclV := validators.NewModelIdOrNameValidator("acl", "loadbalanceracl", ownerProjId)
certV := validators.NewModelIdOrNameValidator("certificate", "loadbalancercertificate", ownerProjId)
tlsCipherPolicyV := validators.NewStringChoicesValidator("tls_cipher_policy", LB_TLS_CIPHER_POLICIES).Default(LB_TLS_CIPHER_POLICY_1_2)
@@ -344,6 +355,9 @@ func (lblis *SLoadbalancerListener) ValidateUpdateData(ctx context.Context, user
return nil, err
}
}
if err := LoadbalancerListenerManager.validateAcl(aclStatusV, aclTypeV, aclV, data); err != nil {
return nil, err
}
{
if backendGroup, ok := backendGroupV.Model.(*SLoadbalancerBackendGroup); ok && backendGroup.LoadbalancerId != lblis.LoadbalancerId {
return nil, httperrors.NewInputParameterError("backend group %s(%s) belongs to loadbalancer %s instead of %s",
+11 -1
View File
@@ -145,7 +145,12 @@ func (self *SNetwork) ValidateDeleteCondition(ctx context.Context) error {
}
func (self *SNetwork) GetTotalNicCount() int {
return self.GetGuestnicsCount() + self.GetGroupNicsCount() + self.GetBaremetalNicsCount() + self.GetReservedNicsCount()
total := self.GetGuestnicsCount() +
self.GetGroupNicsCount() +
self.GetBaremetalNicsCount() +
self.GetReservedNicsCount() +
self.GetLoadbalancerIpsCount()
return total
}
func (self *SNetwork) GetGuestnicsCount() int {
@@ -164,6 +169,10 @@ func (self *SNetwork) GetReservedNicsCount() int {
return ReservedipManager.Query().Equals("network_id", self.Id).Count()
}
func (self *SNetwork) GetLoadbalancerIpsCount() int {
return LoadbalancernetworkManager.Query().Equals("network_id", self.Id).Count()
}
func (self *SNetwork) GetUsedAddresses() map[string]bool {
used := make(map[string]bool)
@@ -818,6 +827,7 @@ func (self *SNetwork) getMoreDetails(extra *jsonutils.JSONDict) *jsonutils.JSOND
extra.Add(jsonutils.JSONFalse, "exit")
}
extra.Add(jsonutils.NewInt(int64(self.getIPRange().AddressCount())), "ports")
extra.Add(jsonutils.NewInt(int64(self.GetTotalNicCount())), "ports_used")
extra.Add(jsonutils.NewInt(int64(self.GetGuestnicsCount())), "vnics")
extra.Add(jsonutils.NewInt(int64(self.GetBaremetalNicsCount())), "bm_vnics")
extra.Add(jsonutils.NewInt(int64(self.GetGroupNicsCount())), "group_vnics")
+17 -22
View File
@@ -9,7 +9,6 @@ import (
"yunion.io/x/jsonutils"
"yunion.io/x/log"
"yunion.io/x/pkg/util/compare"
"yunion.io/x/pkg/utils"
"yunion.io/x/sqlchemy"
"yunion.io/x/onecloud/pkg/cloudcommon/db"
@@ -300,6 +299,7 @@ func (self *SSnapshotManager) CreateSnapshot(ctx context.Context, userCred mccli
return nil, err
}
disk := iDisk.(*SDisk)
storage := disk.GetStorage()
snapshot := &SSnapshot{}
snapshot.SetModelManager(self)
snapshot.ProjectId = userCred.GetProjectId()
@@ -311,6 +311,8 @@ func (self *SSnapshotManager) CreateSnapshot(ctx context.Context, userCred mccli
snapshot.DiskType = disk.DiskType
snapshot.Location = location
snapshot.CreatedBy = createdBy
snapshot.ManagerId = storage.ManagerId
snapshot.CloudregionId = storage.getZone().GetRegion().GetId()
snapshot.Name = name
snapshot.Status = SNAPSHOT_CREATING
err = SnapshotManager.TableSpec().Insert(snapshot)
@@ -345,29 +347,20 @@ func (self *SSnapshot) CustomizeDelete(ctx context.Context, userCred mcclient.To
if self.Status == SNAPSHOT_DELETING {
return fmt.Errorf("Cannot delete snapshot in status %s", self.Status)
}
if self.Status == SNAPSHOT_UNKNOWN {
return self.RealDelete(ctx, userCred)
}
if len(self.ExternalId) == 0 {
if utils.IsInStringArray(self.Status, []string{SNAPSHOT_FAILED}) {
return self.RealDelete(ctx, userCred)
}
if self.CreatedBy == MANUAL {
if !self.FakeDeleted {
return self.FakeDelete()
} else {
_, err := SnapshotManager.GetConvertSnapshot(self)
if err != nil {
return fmt.Errorf("Cannot delete snapshot: %s, disk need at least one of snapshot as backing file", err.Error())
}
return self.StartSnapshotDeleteTask(ctx, userCred, false, "")
}
} else {
return fmt.Errorf("Cannot delete snapshot created by %s", self.CreatedBy)
_, err := SnapshotManager.GetConvertSnapshot(self)
if err != nil {
return fmt.Errorf("Cannot delete snapshot: %s, disk need at least one of snapshot as backing file", err.Error())
}
return self.StartSnapshotDeleteTask(ctx, userCred, false, "")
}
} else {
return self.StartSnapshotDeleteTask(ctx, userCred, false, "")
return fmt.Errorf("Cannot delete snapshot created by %s", self.CreatedBy)
}
return self.StartSnapshotDeleteTask(ctx, userCred, false, "")
}
func (self *SSnapshot) AllowPerformDeleted(ctx context.Context, userCred mcclient.TokenCredential, query jsonutils.JSONObject, data jsonutils.JSONObject) bool {
@@ -467,13 +460,15 @@ func totalSnapshotCount(projectId string) int {
}
// Only sync snapshot status
func (self *SSnapshot) SyncWithCloudSnapshot(userCred mcclient.TokenCredential, ext cloudprovider.ICloudSnapshot, projectId string, projectSync bool) error {
func (self *SSnapshot) SyncWithCloudSnapshot(userCred mcclient.TokenCredential, ext cloudprovider.ICloudSnapshot, projectId string, projectSync bool, region *SCloudregion) error {
_, err := self.GetModelManager().TableSpec().Update(self, func() error {
self.Name = ext.GetName()
self.Status = ext.GetStatus()
self.DiskType = ext.GetDiskType()
if projectSync && len(projectId) > 0 {
self.ProjectId = projectId
}
self.CloudregionId = region.Id
return nil
})
if err != nil {
@@ -482,7 +477,7 @@ func (self *SSnapshot) SyncWithCloudSnapshot(userCred mcclient.TokenCredential,
return err
}
func (manager *SSnapshotManager) newFromCloudSnapshot(userCred mcclient.TokenCredential, extSnapshot cloudprovider.ICloudSnapshot, region *SCloudregion, projectId string) (*SSnapshot, error) {
func (manager *SSnapshotManager) newFromCloudSnapshot(userCred mcclient.TokenCredential, extSnapshot cloudprovider.ICloudSnapshot, region *SCloudregion, projectId string, provider *SCloudprovider) (*SSnapshot, error) {
snapshot := SSnapshot{}
snapshot.SetModelManager(manager)
@@ -500,7 +495,7 @@ func (manager *SSnapshotManager) newFromCloudSnapshot(userCred mcclient.TokenCre
snapshot.DiskType = extSnapshot.GetDiskType()
snapshot.Size = int(extSnapshot.GetSize()) * 1024
snapshot.ManagerId = extSnapshot.GetManagerId()
snapshot.ManagerId = provider.Id
snapshot.CloudregionId = region.Id
snapshot.ProjectId = userCred.GetProjectId()
@@ -555,7 +550,7 @@ func (manager *SSnapshotManager) SyncSnapshots(ctx context.Context, userCred mcc
}
}
for i := 0; i < len(commondb); i += 1 {
err = commondb[i].SyncWithCloudSnapshot(userCred, commonext[i], projectId, projectSync)
err = commondb[i].SyncWithCloudSnapshot(userCred, commonext[i], projectId, projectSync, region)
if err != nil {
syncResult.UpdateError(err)
} else {
@@ -563,7 +558,7 @@ func (manager *SSnapshotManager) SyncSnapshots(ctx context.Context, userCred mcc
}
}
for i := 0; i < len(added); i += 1 {
_, err := manager.newFromCloudSnapshot(userCred, added[i], region, projectId)
_, err := manager.newFromCloudSnapshot(userCred, added[i], region, projectId, provider)
if err != nil {
syncResult.AddError(err)
} else {
+1 -1
View File
@@ -37,7 +37,7 @@ func AttachUsageQuery(
case "vcenter":
q = q.Filter(sqlchemy.Equals(hosts.Field("manager_id"), rangeObjId))
case "schedtag":
aggHosts := SchedtagManager.Query().SubQuery()
aggHosts := HostschedtagManager.Query().SubQuery()
q = q.Join(aggHosts, sqlchemy.AND(
sqlchemy.Equals(hosts.Field("id"), aggHosts.Field("host_id")),
sqlchemy.IsFalse(aggHosts.Field("deleted")))).
-4
View File
@@ -353,10 +353,6 @@ func (manager *SVpcManager) ValidateCreateData(ctx context.Context, userCred mcc
return nil, httperrors.NewInputParameterError("Invalid cloudregion_id")
}
if region.isManaged() {
if region.GetVpcCount() >= MAX_VPC_PER_REGION {
return nil, httperrors.NewNotAcceptableError("Too many vpcs per region")
}
managerStr, _ := data.GetString("manager_id")
if len(managerStr) == 0 {
managerStr, _ = data.GetString("manager")
@@ -353,7 +353,7 @@ func syncZoneHosts(ctx context.Context, provider *models.SCloudprovider, task *C
logSyncFailed(provider, task, msg)
return
}
localHosts, remoteHosts, result := models.HostManager.SyncHosts(ctx, task.UserCred, provider, localZone, hosts)
localHosts, remoteHosts, result := models.HostManager.SyncHosts(ctx, task.UserCred, provider, localZone, hosts, syncRange.ProjectSync)
msg := result.Result()
notes := fmt.Sprintf("SyncHosts for zone %s result: %s", localZone.Name, msg)
log.Infof(notes)
+3
View File
@@ -52,6 +52,7 @@ func (self *GuestCreateTask) OnDiskPreparedFailed(ctx context.Context, obj db.IS
db.OpsLog.LogEvent(guest, db.ACT_ALLOCATE_FAIL, data, self.UserCred)
logclient.AddActionLog(guest, logclient.ACT_ALLOCATE, data, self.UserCred, false)
notifyclient.NotifySystemError(guest.Id, guest.Name, models.VM_DISK_FAILED, data.String())
self.SetStageFailed(ctx, data.String())
}
func (self *GuestCreateTask) OnDiskPrepared(ctx context.Context, obj db.IStandaloneModel, data jsonutils.JSONObject) {
@@ -80,6 +81,7 @@ func (self *GuestCreateTask) OnCdromPreparedFailed(ctx context.Context, obj db.I
db.OpsLog.LogEvent(guest, db.ACT_ALLOCATE_FAIL, data, self.UserCred)
logclient.AddActionLog(guest, logclient.ACT_ALLOCATE, data, self.UserCred, false)
notifyclient.NotifySystemError(guest.Id, guest.Name, models.VM_DISK_FAILED, fmt.Sprintf("cdrom_failed %s", data))
self.SetStageFailed(ctx, fmt.Sprintf("cdrom_failed %s", data))
}
func (self *GuestCreateTask) StartDeployGuest(ctx context.Context, guest *models.SGuest) {
@@ -107,6 +109,7 @@ func (self *GuestCreateTask) OnDeployGuestDescCompleteFailed(ctx context.Context
db.OpsLog.LogEvent(guest, db.ACT_ALLOCATE_FAIL, data, self.UserCred)
logclient.AddActionLog(guest, logclient.ACT_ALLOCATE, data, self.UserCred, false)
notifyclient.NotifySystemError(guest.Id, guest.Name, models.VM_DEPLOY_FAILED, data.String())
self.SetStageFailed(ctx, data.String())
}
func (self *GuestCreateTask) OnAutoStartGuest(ctx context.Context, obj db.IStandaloneModel, data jsonutils.JSONObject) {
+10 -13
View File
@@ -52,32 +52,25 @@ func (self *GuestDiskSnapshotTask) DoDiskSnapshot(ctx context.Context, guest *mo
func (self *GuestDiskSnapshotTask) OnDiskSnapshotComplete(ctx context.Context, guest *models.SGuest, data jsonutils.JSONObject) {
res := data.(*jsonutils.JSONDict)
snapshotId, _ := self.Params.GetString("snapshot_id")
iSnapshot, _ := models.SnapshotManager.FetchById(snapshotId)
snapshot := iSnapshot.(*models.SSnapshot)
if guest.Hypervisor == models.HYPERVISOR_KVM {
location, err := res.GetString("location")
if err != nil {
log.Infof("OnDiskSnapshotComplete called with data no location")
return
}
snapshotId, _ := self.Params.GetString("snapshot_id")
iSnapshot, _ := models.SnapshotManager.FetchById(snapshotId)
snapshot := iSnapshot.(*models.SSnapshot)
models.SnapshotManager.TableSpec().Update(snapshot, func() error {
snapshot.Location = location
snapshot.Status = models.SNAPSHOT_READY
return nil
})
} else {
snapshotId, _ := self.Params.GetString("snapshot_id")
iSnapshot, _ := models.SnapshotManager.FetchById(snapshotId)
snapshot := iSnapshot.(*models.SSnapshot)
extSnapshotId, _ := data.GetString("snapshot_id")
cloudregionId, _ := data.GetString("cloudregion_id")
managerId, _ := data.GetString("manager_id")
models.SnapshotManager.TableSpec().Update(snapshot, func() error {
snapshot.CloudregionId = cloudregionId
snapshot.ExternalId = extSnapshotId
snapshot.Status = models.SNAPSHOT_READY
snapshot.ManagerId = managerId
return nil
})
}
@@ -169,12 +162,16 @@ func (self *SnapshotDeleteTask) deleteExternalSnapshot(ctx context.Context, snap
}
cloudSnapshot, err := cloudRegion.GetISnapshotById(snapshot.ExternalId)
if err != nil {
if err == cloudprovider.ErrNotFound {
return nil
}
log.Errorln(err, cloudSnapshot)
return err
}
cloudSnapshot.Delete()
err = cloudprovider.WaitDeleted(cloudSnapshot, 10*time.Second, 300*time.Second)
return err
if err := cloudSnapshot.Delete(); err != nil {
return err
}
return cloudprovider.WaitDeleted(cloudSnapshot, 10*time.Second, 300*time.Second)
}
func (self *SnapshotDeleteTask) StartReloadDisk(ctx context.Context, snapshot *models.SSnapshot, guest *models.SGuest) {
+1 -1
View File
@@ -272,7 +272,7 @@ func ReportGeneralUsage(userCred mcclient.TokenCredential, rangeObj db.IStandalo
if consts.IsRbacEnabled() {
if policy.PolicyManager.Allow(true, userCred, consts.GetServiceType(),
"usages", policy.PolicyActionGet) == rbacutils.Allow {
"usages", policy.PolicyActionGet) == rbacutils.AdminAllow {
isAdmin = true
}
} else {
+18 -5
View File
@@ -17,25 +17,38 @@ const (
)
func Authenticate(f appsrv.FilterHandler) appsrv.FilterHandler {
return AuthenticateWithDelayDecision(f, false)
}
func AuthenticateWithDelayDecision(f appsrv.FilterHandler, delayDecision bool) appsrv.FilterHandler {
return func(ctx context.Context, w http.ResponseWriter, r *http.Request) {
tokenStr := r.Header.Get(mcclient.AUTH_TOKEN)
if len(tokenStr) == 0 {
httperrors.UnauthorizedError(w, "Unauthorized")
return
log.Errorf("no auth_token found!")
if !delayDecision {
httperrors.UnauthorizedError(w, "Unauthorized")
return
}
}
token, err := Verify(tokenStr)
if err != nil {
log.Errorf("Verify token failed: %s", err)
httperrors.UnauthorizedError(w, "InvalidToken")
return
if !delayDecision {
httperrors.UnauthorizedError(w, "InvalidToken")
return
}
}
ctx = context.WithValue(ctx, AUTH_TOKEN, token)
if token != nil {
ctx = context.WithValue(ctx, AUTH_TOKEN, token)
}
if taskId := r.Header.Get(mcclient.TASK_ID); taskId != "" {
ctx = context.WithValue(ctx, appctx.APP_CONTEXT_KEY_TASK_ID, taskId)
}
if taskNotifyUrl := r.Header.Get(mcclient.TASK_NOTIFY_URL); taskNotifyUrl != "" {
ctx = context.WithValue(ctx, appctx.APP_CONTEXT_KEY_TASK_NOTIFY_URL, taskNotifyUrl)
}
f(ctx, w, r)
}
}
+17
View File
@@ -25,6 +25,15 @@ func NewMonitorManager(keyword, keywordPlural string, columns, adminColumns []st
Keyword: keyword, KeywordPlural: keywordPlural}
}
func NewCloudmonManager(keyword, keywordPlural string, columns, adminColumns []string) ResourceManager {
return ResourceManager{
BaseManager: BaseManager{columns: columns,
adminColumns: adminColumns,
version: "v1",
serviceType: "cloudmon"},
Keyword: keyword, KeywordPlural: keywordPlural}
}
func NewNotifyManager(keyword, keywordPlural string, columns, adminColumns []string) ResourceManager {
return ResourceManager{
BaseManager: BaseManager{columns: columns,
@@ -139,3 +148,11 @@ func NewWebsocketManager(keyword, keywordPlural string, columns, adminColumns []
serviceType: "websocket"},
Keyword: keyword, KeywordPlural: keywordPlural}
}
func NewCloudmetaManager(keyword, keywordPlural string, columns, adminColumns []string) ResourceManager {
return ResourceManager{
BaseManager: BaseManager{columns: columns,
adminColumns: adminColumns,
serviceType: "cloudmeta"},
Keyword: keyword, KeywordPlural: keywordPlural}
}
+13
View File
@@ -0,0 +1,13 @@
package modules
var (
Cloudmeta ResourceManager
)
func init() {
Cloudmeta = NewCloudmetaManager("cloudmeta", "cloudmetas",
[]string{},
[]string{})
register(&Cloudmeta)
}
@@ -0,0 +1,13 @@
package modules
var (
UnderutilizedInstances ResourceManager
)
func init() {
UnderutilizedInstances = NewCloudmonManager("underutilizedinstance", "underutilizedinstances",
[]string{"id", "vm_id", "vm_name", "datetime_str", "vm_cpu", "vm_disk", "vm_memory", "vm_provider", "cpu_usage_threshold", "netio_rx_bps_threshold", "netio_tx_bps_threshold", "stastics_details"},
[]string{})
register(&UnderutilizedInstances)
}
+1 -1
View File
@@ -297,7 +297,7 @@ func (this *ResourceManager) params2Body(s *mcclient.ClientSession, params jsonu
return body
}
func (this *ResourceManager) Create(session *mcclient.ClientSession, params jsonutils.JSONObject) (jsonutils.JSONObject, error) {
func (this *ResourceManager)Create(session *mcclient.ClientSession, params jsonutils.JSONObject) (jsonutils.JSONObject, error) {
return this.CreateInContexts(session, params, nil)
}
+18
View File
@@ -194,3 +194,21 @@ func (o ClusterDeleteNodesOptions) Params() (*jsonutils.JSONDict, error) {
params.Add(nodesArray, "nodes")
return params, nil
}
type ClusterRestartAgentsOptions struct {
ClusterDeleteNodesOptions
All bool `help:"Restart all nodes agent"`
}
func (o ClusterRestartAgentsOptions) Params() (*jsonutils.JSONDict, error) {
params, err := o.ClusterDeleteNodesOptions.Params()
if err != nil {
return nil, err
}
all := jsonutils.JSONFalse
if o.All {
all = jsonutils.JSONTrue
}
params.Add(all, "all")
return params, nil
}
+1 -1
View File
@@ -18,7 +18,7 @@ type ServerListOptions struct {
Gpu *bool `help:"Show gpu servers"`
Secgroup string `help:"Secgroup ID or Name"`
AdminSecgroup string `help:"AdminSecgroup ID or Name"`
Hypervisor string `help:"Show server of hypervisor" choices:"kvm|esxi|container|baremetal|aliyun|azure"`
Hypervisor string `help:"Show server of hypervisor" choices:"kvm|esxi|container|baremetal|aliyun|azure|aws"`
Manager string `help:"Show servers imported from manager"`
Region string `help:"Show servers in cloudregion"`
WithEip *bool `help:"Show Servers with EIP"`
+8 -13
View File
@@ -51,14 +51,6 @@ func (self *SSnapshot) GetStatus() string {
}
}
func (self *SSnapshot) GetManagerId() string {
return self.region.client.providerId
}
func (self *SSnapshot) GetRegionId() string {
return self.region.GetId()
}
func (self *SSnapshot) GetSize() int32 {
return self.SourceDiskSize
}
@@ -168,13 +160,16 @@ func (self *SRegion) GetSnapshots(instanceId string, diskId string, snapshotName
}
func (self *SRegion) GetISnapshotById(snapshotId string) (cloudprovider.ICloudSnapshot, error) {
if snapshots, total, err := self.GetSnapshots("", "", "", []string{snapshotId}, 0, 1); err != nil {
snapshots, total, err := self.GetSnapshots("", "", "", []string{snapshotId}, 0, 1)
if err != nil {
return nil, err
} else if total != 1 {
return nil, cloudprovider.ErrNotFound
} else {
return &snapshots[0], nil
}
if total == 0 {
return nil, cloudprovider.ErrNotFound
} else if total > 1 {
return nil, cloudprovider.ErrDuplicateId
}
return &snapshots[0], nil
}
func (self *SRegion) DeleteSnapshot(snapshotId string) error {
+1 -1
View File
@@ -1,5 +1,5 @@
package ansible
const (
PUBLIC_CLOUD_ANSIBLE_USER = "yunionroot"
PUBLIC_CLOUD_ANSIBLE_USER = "cloudroot"
)
+2 -1
View File
@@ -1,7 +1,7 @@
package aws
import (
"github.com/coredns/coredns/plugin/pkg/log"
"yunion.io/x/log"
"yunion.io/x/onecloud/pkg/cloudprovider"
"yunion.io/x/onecloud/pkg/compute/models"
@@ -33,6 +33,7 @@ func NewAwsClient(providerId string, providerName string, accessUrl string, acce
client := SAwsClient{providerId: providerId, providerName: providerName, accessUrl: accessUrl, accessKey: accessKey, secret: secret}
err := client.fetchRegions()
if err != nil {
log.Debugf("NewAwsClient %s", err.Error())
return nil, err
}
return &client, nil
+51 -3
View File
@@ -28,7 +28,7 @@ type SDisk struct {
DiskId string // VolumeId
DiskName string // Tag Name
Size int // Size
Size int // Size GB
Category string // VolumeType
Type string // system | data
Status string // State
@@ -339,6 +339,7 @@ func (self *SRegion) DeleteDisk(diskId string) error {
}
params.SetVolumeId(diskId)
log.Debugf("DeleteDisk with params: %s", params.String())
_, err = self.ec2Client.DeleteVolume(params)
return err
}
@@ -366,8 +367,49 @@ func (self *SRegion) resizeDisk(diskId string, size int64) error {
}
func (self *SRegion) resetDisk(diskId, snapshotId string) error {
// aws貌似不支持直接重置
return cloudprovider.ErrNotImplemented
// 这里实际是回滚快照
disk, err := self.GetDisk(diskId)
if err != nil {
log.Debugf("resetDisk %s:%s", diskId, err.Error())
return err
}
params := &ec2.CreateVolumeInput{}
params.SetSnapshotId(snapshotId)
params.SetSize(int64(disk.Size))
params.SetVolumeType(disk.Category)
params.SetAvailabilityZone(disk.ZoneId)
tags, _ := disk.Tags.GetTagSpecifications()
params.SetTagSpecifications([]*ec2.TagSpecification{tags})
ret, err := self.ec2Client.CreateVolume(params)
if err != nil {
log.Debugf("resetDisk %s: %s", params.String(), err.Error())
return err
}
// detach disk
if disk.Status == ec2.VolumeStateInUse {
err := self.DetachDisk(disk.InstanceId, diskId)
if err != nil {
log.Debugf("resetDisk %s %s: %s", disk.InstanceId, diskId, err.Error())
return err
}
err = self.ec2Client.WaitUntilVolumeAvailable(&ec2.DescribeVolumesInput{VolumeIds: []*string{&diskId}})
if err != nil {
log.Debugf("resetDisk :%s", err.Error())
return err
}
}
err = self.AttachDisk(disk.InstanceId, *ret.VolumeId, disk.Device)
if err != nil {
log.Debugf("resetDisk %s %s %s: %s", disk.InstanceId, *ret.VolumeId, disk.Device, err.Error())
return err
}
// 绑定成功后删除原磁盘
return self.DeleteDisk(diskId)
}
func (self *SRegion) CreateDisk(zoneId string, category string, name string, sizeGb int, snapshotId string, desc string) (string, error) {
@@ -390,5 +432,11 @@ func (self *SRegion) CreateDisk(zoneId string, category string, name string, siz
return "", err
}
paramsWait := &ec2.DescribeVolumesInput{}
paramsWait.SetVolumeIds([]*string{ret.VolumeId})
err = self.ec2Client.WaitUntilVolumeAvailable(paramsWait)
if err != nil {
return "", err
}
return StrVal(ret.VolumeId), nil
}
+11
View File
@@ -2,6 +2,7 @@ package aws
import (
"fmt"
"strings"
"github.com/aws/aws-sdk-go/service/ec2"
"yunion.io/x/jsonutils"
@@ -189,6 +190,8 @@ func (self *SRegion) GetImageByName(name string) (*SImage, error) {
if len(images) == 0 {
return nil, cloudprovider.ErrNotFound
}
log.Debugf("%d image found match name %", len(images), name)
return &images[0], nil
}
@@ -230,8 +233,16 @@ func (self *SRegion) GetImages(status ImageStatusType, owner ImageOwnerType, ima
if len(imageId) > 0 {
params.SetImageIds(ConvertedList(imageId))
}
if len(filters) > 0 {
params.SetFilters(filters)
}
ret, err := self.ec2Client.DescribeImages(params)
if err != nil {
if strings.Contains(err.Error(), ".NotFound") {
return nil, 0, cloudprovider.ErrNotFound
}
return nil, 0, err
}
+32 -17
View File
@@ -6,8 +6,8 @@ import (
"context"
"github.com/aws/aws-sdk-go/service/ec2"
"github.com/coredns/coredns/plugin/pkg/log"
"yunion.io/x/jsonutils"
"yunion.io/x/log"
"yunion.io/x/onecloud/pkg/cloudprovider"
"yunion.io/x/onecloud/pkg/compute/models"
"yunion.io/x/pkg/util/osprofile"
@@ -274,23 +274,26 @@ func (self *SInstance) SyncSecurityGroup(secgroupId string, name string, rules [
if vpc, err := self.getVpc(); err != nil {
return err
} else if len(secgroupId) == 0 {
for index, secgrpId := range self.SecurityGroupIds.SecurityGroupId {
if err := vpc.revokeSecurityGroup(secgrpId, self.InstanceId, index == 0); err != nil {
return err
}
}
// todo : 这里应该有问题。aws不能直接删除安全组。且至少选择一个安全组
// for index, secgrpId := range self.SecurityGroupIds.SecurityGroupId {
// if err := vpc.revokeSecurityGroup(secgrpId, self.InstanceId, index == 0); err != nil {
// return err
// }
// }
return nil
} else if secgrpId, err := vpc.SyncSecurityGroup(secgroupId, name, rules); err != nil {
return err
} else if err := vpc.assignSecurityGroup(secgrpId, self.InstanceId); err != nil {
return err
} else {
for _, secgroupId := range self.SecurityGroupIds.SecurityGroupId {
if secgroupId != secgrpId {
if err := vpc.revokeSecurityGroup(secgroupId, self.InstanceId, false); err != nil {
return err
}
}
}
// todo : 这里应该有问题。aws不能直接删除安全组。且至少选择一个安全组
// for _, secgroupId := range self.SecurityGroupIds.SecurityGroupId {
// if secgroupId != secgrpId {
// if err := vpc.revokeSecurityGroup(secgroupId, self.InstanceId, false); err != nil {
// return err
// }
// }
// }
self.SecurityGroupIds.SecurityGroupId = []string{secgrpId}
}
return nil
@@ -375,12 +378,18 @@ func (self *SInstance) GetVNCInfo() (jsonutils.JSONObject, error) {
}
func (self *SInstance) AttachDisk(ctx context.Context, diskId string) error {
// todo:bugfix . self.DeviceNames => self.GetDeviceNames()
name, err := NextDeviceName(self.DeviceNames)
if err != nil {
return err
}
return self.host.zone.region.AttachDisk(self.InstanceId, diskId, name)
err = self.host.zone.region.AttachDisk(self.InstanceId, diskId, name)
if err != nil {
return err
}
self.DeviceNames = append(self.DeviceNames, name)
return nil
}
func (self *SInstance) DetachDisk(ctx context.Context, diskId string) error {
@@ -412,11 +421,9 @@ func (self *SRegion) GetInstances(zoneId string, ids []string, offset int, limit
log.Errorf("GetInstances fail %s", err)
return nil, 0, err
}
instances := []SInstance{}
for _, reservation := range res.Reservations {
for _, instance := range reservation.Instances {
log.Debugf("GetInstances %s", instance.String())
if err := FillZero(instance); err != nil {
return nil, 0, err
}
@@ -465,8 +472,16 @@ func (self *SRegion) GetInstances(zoneId string, ids []string, offset int, limit
productCodes = append(productCodes, *p.ProductCodeId)
}
szone, err := self.getZoneById(*instance.Placement.AvailabilityZone)
if err != nil {
return nil, 0, err
}
host := szone.getHost()
sinstance := SInstance{
RegionId: self.RegionId,
host: host,
ZoneId: *instance.Placement.AvailabilityZone,
InstanceId: *instance.InstanceId,
ImageId: *instance.ImageId,
+23
View File
@@ -0,0 +1,23 @@
package aws
var LatitudeAndLongitude = map[string]map[string]float32{
"ap-south-1": {"latitude": 19.0759837, "longitude": 72.8776559},
"ap-northeast-3": {"latitude": 34.6937378, "longitude": 135.5021651},
"us-east-1": {"latitude": 37.4315734, "longitude": -78.6568942},
"us-east-2": {"latitude": 40.4172871, "longitude": -82.90712300000001},
"ap-southeast-2": {"latitude": -33.8688197, "longitude": 151.2092955},
"cn-northwest-1": {"latitude": 37.198731, "longitude": 106.1580937},
"eu-west-1": {"latitude": 53.41291, "longitude": -8.24389},
"eu-central-1": {"latitude": 50.1109221, "longitude": 8.6821267},
"sa-east-1": {"latitude": -23.5505199, "longitude": -46.63330939999999},
"ap-southeast-1": {"latitude": 1.352083, "longitude": 103.819836},
"ca-central-1": {"latitude": 56.130366, "longitude": -106.346771},
"ap-northeast-2": {"latitude": 37.566535, "longitude": 126.9779692},
"us-west-2": {"latitude": 43.8041334, "longitude": -120.5542012},
"us-gov-west-1": {"latitude": 37.09024, "longitude": -95.712891},
"us-west-1": {"latitude": 38.8375215, "longitude": -120.8958242},
"cn-north-1": {"latitude": 39.90419989999999, "longitude": 116.4073963},
"ap-northeast-1": {"latitude": 35.7090259, "longitude": 139.7319925},
"eu-west-2": {"latitude": 51.5073509, "longitude": -0.1277583},
"eu-west-3": {"latitude": 48.856614, "longitude": 2.3522219},
}
+76 -14
View File
@@ -2,6 +2,7 @@ package aws
import (
"fmt"
sdk "github.com/aws/aws-sdk-go/aws"
"github.com/aws/aws-sdk-go/aws/credentials"
"github.com/aws/aws-sdk-go/aws/session"
@@ -14,6 +15,28 @@ import (
"yunion.io/x/onecloud/pkg/compute/models"
)
var RegionLocations map[string]string = map[string]string{
"us-east-2": "美国东部(俄亥俄州)",
"us-east-1": "美国东部(弗吉尼亚北部)",
"us-west-1": "美国西部(加利福尼亚北部)",
"us-west-2": "美国西部(俄勒冈)",
"ap-south-1": "亚太地区(孟买)",
"ap-northeast-2": "亚太区域(首尔)",
"ap-northeast-3": "亚太区域(大阪)",
"ap-southeast-1": "亚太区域(新加坡)",
"ap-southeast-2": "亚太区域(悉尼)",
"ap-northeast-1": "亚太区域(东京)",
"ca-central-1": "加拿大(中部)",
"cn-north-1": "中国(北京)",
"cn-northwest-1": "中国(宁夏)",
"eu-central-1": "欧洲(法兰克福)",
"eu-west-1": "欧洲(爱尔兰)",
"eu-west-2": "欧洲(伦敦)",
"eu-west-3": "欧洲(巴黎)",
"sa-east-1": "南美洲(圣保罗)",
"us-gov-west-1": "AWS GovCloud(美国)",
}
type SRegion struct {
client *SAwsClient
ec2Client *ec2.EC2
@@ -36,12 +59,16 @@ func (self *SRegion) GetClient() *SAwsClient {
return self.client
}
func (self *SRegion) getAwsSession() (*session.Session, error) {
return session.NewSession(&sdk.Config{
Region: sdk.String(self.RegionId),
Credentials: credentials.NewStaticCredentials(self.client.accessKey, self.client.secret, ""),
})
}
func (self *SRegion) getEc2Client() (*ec2.EC2, error) {
if self.ec2Client == nil {
s, err := session.NewSession(&sdk.Config{
Region: sdk.String(self.RegionId),
Credentials: credentials.NewStaticCredentials(self.client.accessKey, self.client.secret, ""),
})
s, err := self.getAwsSession()
if err != nil {
return nil, err
@@ -56,10 +83,7 @@ func (self *SRegion) getEc2Client() (*ec2.EC2, error) {
func (self *SRegion) getIamClient() (*iam.IAM, error) {
if self.iamClient == nil {
s, err := session.NewSession(&sdk.Config{
Region: sdk.String(self.RegionId),
Credentials: credentials.NewStaticCredentials(self.client.accessKey, self.client.secret, ""),
})
s, err := self.getAwsSession()
if err != nil {
return nil, err
@@ -73,10 +97,7 @@ func (self *SRegion) getIamClient() (*iam.IAM, error) {
func (self *SRegion) getS3Client() (*s3.S3, error) {
if self.s3Client == nil {
s, err := session.NewSession(&sdk.Config{
Region: sdk.String(self.RegionId),
Credentials: credentials.NewStaticCredentials(self.client.accessKey, self.client.secret, ""),
})
s, err := self.getAwsSession()
if err != nil {
return nil, err
@@ -164,6 +185,10 @@ func (self *SRegion) GetId() string {
}
func (self *SRegion) GetName() string {
if localName, ok := RegionLocations[self.RegionId]; ok {
return fmt.Sprintf("%s %s", CLOUD_PROVIDER_AWS_CN, localName)
}
return fmt.Sprintf("%s %s", CLOUD_PROVIDER_AWS_CN, self.RegionId)
}
@@ -188,11 +213,27 @@ func (self *SRegion) GetMetadata() *jsonutils.JSONDict {
}
func (self *SRegion) GetLatitude() float32 {
return 0.0
if data, ok := LatitudeAndLongitude[self.RegionId]; !ok {
log.Debugf("Region %s not found in LatitudeAndLongitude", self.RegionId)
return 0.0
} else if lat, ok := data["latitude"]; !ok {
log.Debugf("Region %s's latitude not found in LatitudeAndLongitude", self.RegionId)
return 0.0
} else {
return lat
}
}
func (self *SRegion) GetLongitude() float32 {
return 0.0
if data, ok := LatitudeAndLongitude[self.RegionId]; !ok {
log.Debugf("Region %s not found in LatitudeAndLongitude", self.RegionId)
return 0.0
} else if lat, ok := data["longitude"]; !ok {
log.Debugf("Region %s's latitude not found in LatitudeAndLongitude", self.RegionId)
return 0.0
} else {
return lat
}
}
func (self *SRegion) GetIZones() ([]cloudprovider.ICloudZone, error) {
@@ -320,11 +361,32 @@ func (self *SRegion) GetIStoragecacheById(id string) (cloudprovider.ICloudStorag
}
func (self *SRegion) CreateIVpc(name string, desc string, cidr string) (cloudprovider.ICloudVpc, error) {
tagspec := TagSpec{ResourceType: "vpc"}
if len(name) > 0 {
tagspec.SetNameTag(name)
}
if len(desc) > 0 {
tagspec.SetDescTag(desc)
}
spec, err := tagspec.GetTagSpecifications()
if err != nil {
return nil, err
}
// start create vpc
vpc, err := self.ec2Client.CreateVpc(&ec2.CreateVpcInput{CidrBlock: &cidr})
if err != nil {
return nil, err
}
tagsParams := &ec2.CreateTagsInput{Resources: []*string{vpc.Vpc.VpcId}, Tags: spec.Tags}
_, err = self.ec2Client.CreateTags(tagsParams)
if err != nil {
log.Debugf("CreateIVpc add tag failed %s", err.Error())
}
err = self.fetchInfrastructure()
if err != nil {
return nil, err
+49 -41
View File
@@ -27,7 +27,7 @@ type SSecurityGroup struct {
VpcId string
SecurityGroupId string
Description string
SecurityGroupName string
SecurityGroupName string //对应tag中的name标签
Permissions []secrules.SecurityRule
Tags Tags
@@ -117,22 +117,22 @@ func (self *SRegion) addSecurityGroupRule(secGrpId string, rule *secrules.Securi
params := &ec2.AuthorizeSecurityGroupIngressInput{}
params.SetGroupId(secGrpId)
params.SetIpPermissions(ipPermissions)
_, err := self.ec2Client.AuthorizeSecurityGroupIngress(params)
if err != nil {
return err
}
_, err = self.ec2Client.AuthorizeSecurityGroupIngress(params)
}
if rule.Direction == secrules.SecurityRuleEgress {
params := &ec2.AuthorizeSecurityGroupEgressInput{}
params.SetGroupId(secGrpId)
params.SetIpPermissions(ipPermissions)
_, err := self.ec2Client.AuthorizeSecurityGroupEgress(params)
if err != nil {
return err
}
_, err = self.ec2Client.AuthorizeSecurityGroupEgress(params)
}
return nil
if err != nil && strings.Contains(err.Error(), "InvalidPermission.Duplicate") {
log.Debugf("addSecurityGroupRule %s %s", rule.Direction, err.Error())
return nil
}
return err
}
func (self *SRegion) delSecurityGroupRule(secGrpId string, rule *secrules.SecurityRule) error {
@@ -145,20 +145,19 @@ func (self *SRegion) delSecurityGroupRule(secGrpId string, rule *secrules.Securi
params := &ec2.RevokeSecurityGroupIngressInput{}
params.SetGroupId(secGrpId)
params.SetIpPermissions(ipPermissions)
_, err := self.ec2Client.RevokeSecurityGroupIngress(params)
if err != nil {
return err
}
_, err = self.ec2Client.RevokeSecurityGroupIngress(params)
}
if rule.Direction == secrules.SecurityRuleEgress {
params := &ec2.RevokeSecurityGroupEgressInput{}
params.SetGroupId(secGrpId)
params.SetIpPermissions(ipPermissions)
_, err := self.ec2Client.RevokeSecurityGroupEgress(params)
if err != nil {
return err
}
_, err = self.ec2Client.RevokeSecurityGroupEgress(params)
}
if err != nil {
log.Debugf("delSecurityGroupRule %s %s", rule.Direction, err.Error())
return err
}
return nil
}
@@ -198,8 +197,10 @@ func (self *SRegion) updateSecurityGroupRuleDescription(secGrpId string, rule *s
func (self *SRegion) createSecurityGroup(vpcId string, name string, secgroupIdTag string, desc string) (string, error) {
params := &ec2.CreateSecurityGroupInput{}
params.SetVpcId(vpcId)
// 这里的描述aws 上层代码拼接的描述。并非用户提交的描述,用户描述放置在Yunion本地数据库中。)
params.SetDescription(desc)
params.SetGroupName(name)
// 这里使用id作为组名。原因name容易重名、另外有可能包含中文,aws不支持中文
params.SetGroupName(secgroupIdTag)
group, err := self.ec2Client.CreateSecurityGroup(params)
if err != nil {
@@ -208,6 +209,8 @@ func (self *SRegion) createSecurityGroup(vpcId string, name string, secgroupIdTa
tagspec := TagSpec{ResourceType: "security-group"}
tagspec.SetTag("id", secgroupIdTag)
tagspec.SetNameTag(name)
tagspec.SetDescTag(desc)
tags, _ := tagspec.GetTagSpecifications()
tagParams := &ec2.CreateTagsInput{}
tagParams.SetResources([]*string{group.GroupId})
@@ -308,6 +311,9 @@ func (self *SRegion) modifySecurityGroup(secGrpId string, name string, desc stri
}
func (self *SRegion) syncSecgroupRules(secgroupId string, rules []secrules.SecurityRule) error {
var DeleteRules []secrules.SecurityRule
var AddRules []secrules.SecurityRule
if secgroup, err := self.GetSecurityGroupDetails(secgroupId); err != nil {
return err
} else {
@@ -315,6 +321,9 @@ func (self *SRegion) syncSecgroupRules(secgroupId string, rules []secrules.Secur
sort.Sort(secrules.SecurityRuleSet(rules))
sort.Sort(secrules.SecurityRuleSet(secgroup.Permissions))
log.Debugf("local security rules %s", rules)
log.Debugf("remote security rules %s", secgroup.Permissions)
i, j := 0, 0
for i < len(rules) || j < len(secgroup.Permissions) {
if i < len(rules) && j < len(secgroup.Permissions) {
@@ -322,42 +331,41 @@ func (self *SRegion) syncSecgroupRules(secgroupId string, rules []secrules.Secur
ruleStr := rules[i].String()
cmp := strings.Compare(permissionStr, ruleStr)
if cmp == 0 {
if secgroup.Permissions[j].Description != rules[i].Description {
if err := self.updateSecurityGroupRuleDescription(secgroupId, &rules[i]); err != nil {
log.Errorf("updateSecurityGroupRuleDescription error %v", rules[i])
return err
}
}
DeleteRules = append(DeleteRules, secgroup.Permissions[j])
AddRules = append(AddRules, rules[i])
i += 1
j += 1
} else if cmp > 0 {
if err := self.delSecurityGroupRule(secgroupId, &secgroup.Permissions[j]); err != nil {
log.Errorf("delSecurityGroupRule error %v", secgroup.Permissions[j])
return err
}
DeleteRules = append(DeleteRules, secgroup.Permissions[j])
j += 1
} else {
if err := self.addSecurityGroupRules(secgroupId, &rules[i]); err != nil {
log.Errorf("addSecurityGroupRule error %v", rules[i])
return err
}
AddRules = append(AddRules, rules[i])
i += 1
}
} else if i >= len(rules) {
if err := self.delSecurityGroupRule(secgroupId, &secgroup.Permissions[j]); err != nil {
log.Errorf("delSecurityGroupRule error %v", secgroup.Permissions[j])
return err
}
DeleteRules = append(DeleteRules, secgroup.Permissions[j])
j += 1
} else if j >= len(secgroup.Permissions) {
if err := self.addSecurityGroupRules(secgroupId, &rules[i]); err != nil {
log.Errorf("addSecurityGroupRule error %v", rules[i])
return err
}
AddRules = append(AddRules, rules[i])
i += 1
}
}
}
for _, r := range DeleteRules {
if err := self.delSecurityGroupRule(secgroupId, &r); err != nil {
log.Errorf("delSecurityGroupRule %v error: %s", r, err.Error())
return err
}
}
for _, r := range AddRules {
if err := self.addSecurityGroupRules(secgroupId, &r); err != nil {
log.Errorf("addSecurityGroupRule %v error: %s", r, err.Error())
return err
}
}
return nil
}
+14 -6
View File
@@ -3,7 +3,9 @@ package aws
import (
"fmt"
"github.com/aws/aws-sdk-go/service/ec2"
"strings"
"yunion.io/x/jsonutils"
"yunion.io/x/log"
"yunion.io/x/onecloud/pkg/cloudprovider"
"yunion.io/x/onecloud/pkg/compute/models"
)
@@ -11,9 +13,9 @@ import (
type SnapshotStatusType string
const (
SnapshotStatusAccomplished SnapshotStatusType = "accomplished"
SnapshotStatusProgress SnapshotStatusType = "progressing"
SnapshotStatusFailed SnapshotStatusType = "failed"
SnapshotStatusAccomplished SnapshotStatusType = "completed"
SnapshotStatusProgress SnapshotStatusType = "pending"
SnapshotStatusFailed SnapshotStatusType = "error"
)
type SSnapshot struct {
@@ -89,10 +91,11 @@ func (self *SSnapshot) GetDiskId() string {
}
func (self *SSnapshot) Delete() error {
panic("implement me")
return self.region.DeleteSnapshot(self.SnapshotId)
}
func (self *SSnapshot) GetRegionId() string {
// 这里特别注意:aws没有有uuid形式的region id
return self.region.GetId()
}
@@ -124,6 +127,10 @@ func (self *SRegion) GetSnapshots(instanceId string, diskId string, snapshotName
ret, err := self.ec2Client.DescribeSnapshots(params)
if err != nil {
if strings.Contains(err.Error(), "InvalidSnapshot.NotFound") {
return nil, 0, cloudprovider.ErrNotFound
}
return nil, 0, err
}
@@ -176,8 +183,9 @@ func (self *SRegion) CreateSnapshot(diskId, name, desc string) (string, error) {
}
params.SetDescription(desc)
_, err := self.ec2Client.CreateSnapshot(params)
return "", err
log.Debugf("CreateSnapshots with params %s", params)
ret, err := self.ec2Client.CreateSnapshot(params)
return StrVal(ret.SnapshotId), err
}
func (self *SRegion) DeleteSnapshot(snapshotId string) error {
+16 -13
View File
@@ -1,12 +1,11 @@
package aws
import (
"bytes"
"fmt"
"github.com/aws/aws-sdk-go/service/ec2"
"github.com/aws/aws-sdk-go/service/iam"
"github.com/aws/aws-sdk-go/service/s3"
"io/ioutil"
"github.com/aws/aws-sdk-go/service/s3/s3manager"
"strings"
"time"
"yunion.io/x/jsonutils"
@@ -148,17 +147,20 @@ func (self *SStoragecache) uploadImage(userCred mcclient.TokenCredential, imageI
return "", err
}
s3Client, err := self.region.getS3Client()
if err != nil {
return "", nil
// uploader to aws s3
input := &s3manager.UploadInput{
Bucket: &bucketName,
Key: &imageId,
Body: reader,
}
// 内存?
f, err := ioutil.ReadAll(reader)
params := &s3.PutObjectInput{}
params.SetBucket(bucketName)
params.SetKey(imageId)
params.SetBody(bytes.NewReader(f))
_, err = s3Client.PutObject(params)
awsSession, err := self.region.getAwsSession()
if err != nil {
log.Debugf("uploadImage %s", err.Error())
return "", fmt.Errorf("get aws session failed")
}
uploader := s3manager.NewUploader(awsSession)
_, err = uploader.Upload(input)
if err != nil {
return "", nil
}
@@ -183,6 +185,7 @@ func (self *SStoragecache) uploadImage(userCred mcclient.TokenCredential, imageI
imageName = fmt.Sprintf("%s-%d", imageBaseName, nameIdx)
nameIdx += 1
log.Debugf("uploadImage Match remote name %s", imageName)
}
task, err := self.region.ImportImage(imageName, osArch, osType, osDist, diskFormat, bucketName, imageId)
@@ -194,6 +197,7 @@ func (self *SStoragecache) uploadImage(userCred mcclient.TokenCredential, imageI
// todo:// 等待镜像导入完成
for i := 1; i < 120; i++ {
time.Sleep(2 * time.Minute)
ret, err := self.region.ec2Client.DescribeImportImageTasks(&ec2.DescribeImportImageTasksInput{ImportTaskIds: []*string{&task.TaskId}})
if err != nil {
return "", err
@@ -210,7 +214,6 @@ func (self *SStoragecache) uploadImage(userCred mcclient.TokenCredential, imageI
return *item.ImageId, nil
}
}
time.Sleep(1 * time.Minute)
}
return task.ImageId, fmt.Errorf("uploadImage uncompleted: %s", task)
+9 -1
View File
@@ -4,6 +4,7 @@ import (
"fmt"
"net"
"reflect"
"regexp"
"strings"
"yunion.io/x/jsonutils"
@@ -295,6 +296,8 @@ func AwsIpPermissionToYunion(direction secrules.TSecurityRuleDirection, p ec2.Ip
return rules, nil
}
// YunionSecRuleToAws 不能保证无损转换
// 规则描述如果包含中文等字符,将被丢弃掉
func YunionSecRuleToAws(rule secrules.SecurityRule) ([]*ec2.IpPermission, error) {
if rule.Action == secrules.SecurityRuleDeny {
return nil, fmt.Errorf("YunionSecRuleToAws ignored aws not supported deny rule")
@@ -304,8 +307,13 @@ func YunionSecRuleToAws(rule secrules.SecurityRule) ([]*ec2.IpPermission, error)
if iprange == "<nil>" {
return nil, fmt.Errorf("YunionSecRuleToAws ignored ipnet should not be empty")
}
description := ""
if match, err := regexp.MatchString("^[\\sa-zA-Z0-9. _:/()#,@\\]\\[+=&;{}!$*-]+$", rule.Description); err == nil && match {
description = rule.Description
}
ipranges := []*ec2.IpRange{}
ipranges = append(ipranges, &ec2.IpRange{CidrIp: &iprange, Description: &rule.Description})
ipranges = append(ipranges, &ec2.IpRange{CidrIp: &iprange, Description: &description})
portranges := yunionPortRangeToAws(rule)
protocol := yunionProtocolToAws(rule)
+31 -15
View File
@@ -108,19 +108,7 @@ func (self *SVpc) GetManagerId() string {
}
func (self *SVpc) Delete() error {
err := self.fetchSecurityGroups()
if err != nil {
log.Errorf("fetchSecurityGroup for VPC delete fail %s", err)
return err
}
for i := 0; i < len(self.secgroups); i += 1 {
secgroup := self.secgroups[i].(*SSecurityGroup)
err := self.region.deleteSecurityGroup(secgroup.SecurityGroupId)
if err != nil {
log.Errorf("deleteSecurityGroup for VPC delete fail %s", err)
return err
}
}
// 删除vpc会同步删除关联的安全组
return self.region.DeleteVpc(self.VpcId)
}
@@ -240,19 +228,44 @@ func (self *SRegion) getVpc(vpcId string) (*SVpc, error) {
}
func (self *SRegion) revokeSecurityGroup(secgroupId, instanceId string, keep bool) error {
// todo : keep ? 直接使用assignSecurityGroup 即可?
return nil
}
func (self *SRegion) assignSecurityGroup(secgroupId, instanceId string) error {
instance, err := self.GetInstance(instanceId)
if err != nil {
return err
}
for _, eth := range instance.NetworkInterfaces.NetworkInterface {
params := &ec2.ModifyNetworkInterfaceAttributeInput{}
params.SetNetworkInterfaceId(eth.NetworkInterfaceId)
params.SetGroups([]*string{&secgroupId})
_, err := self.ec2Client.ModifyNetworkInterfaceAttribute(params)
if err != nil {
return err
}
}
return nil
}
func (self *SRegion) deleteSecurityGroup(secGrpId string) error {
return nil
params := &ec2.DeleteSecurityGroupInput{}
params.SetGroupId(secGrpId)
_, err := self.ec2Client.DeleteSecurityGroup(params)
return err
}
func (self *SRegion) DeleteVpc(vpcId string) error {
return nil
params := &ec2.DeleteVpcInput{}
params.SetVpcId(vpcId)
_, err := self.ec2Client.DeleteVpc(params)
return err
}
func (self *SRegion) GetVpcs(vpcId []string, offset int, limit int) ([]SVpc, int, error) {
@@ -262,6 +275,9 @@ func (self *SRegion) GetVpcs(vpcId []string, offset int, limit int) ([]SVpc, int
}
ret, err := self.ec2Client.DescribeVpcs(params)
if err != nil {
if strings.Contains(err.Error(), "InvalidVpcID.NotFound") {
return nil, 0, cloudprovider.ErrNotFound
}
return nil, 0, err
}
+115 -28
View File
@@ -106,8 +106,10 @@ func (self *SAzureClient) getDefaultClient() (*autorest.Client, error) {
return nil, err
}
client.Authorizer = authorizer
// client.RequestInspector = LogRequest()
// client.ResponseInspector = LogResponse()
if DEBUG {
client.RequestInspector = LogRequest()
client.ResponseInspector = LogResponse()
}
return &client, nil
}
@@ -116,7 +118,7 @@ func (self *SAzureClient) jsonRequest(method, url string, body string) (jsonutil
if err != nil {
return nil, err
}
return jsonRequest(cli, method, self.domain, url, body)
return jsonRequest(cli, method, self.domain, url, self.subscriptionId, body)
}
func (self *SAzureClient) Get(resourceId string, params []string, retVal interface{}) error {
@@ -131,7 +133,7 @@ func (self *SAzureClient) Get(resourceId string, params []string, retVal interfa
if err != nil {
return err
}
body, err := jsonRequest(cli, "GET", self.domain, path, "")
body, err := jsonRequest(cli, "GET", self.domain, path, self.subscriptionId, "")
if err != nil {
return err
}
@@ -151,7 +153,7 @@ func (self *SAzureClient) ListVmSizes(location string) (jsonutils.JSONObject, er
return nil, fmt.Errorf("need subscription id")
}
url := fmt.Sprintf("/subscriptions/%s/providers/Microsoft.Compute/locations/%s/vmSizes", self.subscriptionId, location)
return jsonRequest(cli, "GET", self.domain, url, "")
return jsonRequest(cli, "GET", self.domain, url, self.subscriptionId, "")
}
func (self *SAzureClient) ListClassicDisks() (jsonutils.JSONObject, error) {
@@ -163,7 +165,7 @@ func (self *SAzureClient) ListClassicDisks() (jsonutils.JSONObject, error) {
return nil, fmt.Errorf("need subscription id")
}
url := fmt.Sprintf("/subscriptions/%s/services/disks", self.subscriptionId)
return jsonRequest(cli, "GET", self.domain, url, "")
return jsonRequest(cli, "GET", self.domain, url, self.subscriptionId, "")
}
func (self *SAzureClient) ListAll(resourceType string, retVal interface{}) error {
@@ -178,7 +180,7 @@ func (self *SAzureClient) ListAll(resourceType string, retVal interface{}) error
if len(resourceType) > 0 {
url += fmt.Sprintf("/providers/%s", resourceType)
}
body, err := jsonRequest(cli, "GET", self.domain, url, "")
body, err := jsonRequest(cli, "GET", self.domain, url, self.subscriptionId, "")
if err != nil {
return err
}
@@ -193,7 +195,7 @@ func (self *SAzureClient) ListSubscriptions() (jsonutils.JSONObject, error) {
if err != nil {
return nil, err
}
return jsonRequest(cli, "GET", self.domain, "/subscriptions", "")
return jsonRequest(cli, "GET", self.domain, "/subscriptions", self.subscriptionId, "")
}
func (self *SAzureClient) List(golbalResource string, retVal interface{}) error {
@@ -208,7 +210,7 @@ func (self *SAzureClient) List(golbalResource string, retVal interface{}) error
if len(self.subscriptionId) > 0 && len(golbalResource) > 0 {
url += fmt.Sprintf("/%s", golbalResource)
}
body, err := jsonRequest(cli, "GET", self.domain, url, "")
body, err := jsonRequest(cli, "GET", self.domain, url, self.subscriptionId, "")
if err != nil {
return err
}
@@ -224,7 +226,7 @@ func (self *SAzureClient) ListByTypeWithResourceGroup(resourceGroupName string,
return fmt.Errorf("Missing subscription Info")
}
url := fmt.Sprintf("/subscriptions/%s/resourceGroups/%s/providers/%s", self.subscriptionId, resourceGroupName, Type)
body, err := jsonRequest(cli, "GET", self.domain, url, "")
body, err := jsonRequest(cli, "GET", self.domain, url, self.subscriptionId, "")
if err != nil {
return err
}
@@ -236,7 +238,7 @@ func (self *SAzureClient) Delete(resourceId string) error {
if err != nil {
return err
}
_, err = jsonRequest(cli, "DELETE", self.domain, resourceId, "")
_, err = jsonRequest(cli, "DELETE", self.domain, resourceId, self.subscriptionId, "")
return err
}
@@ -246,7 +248,7 @@ func (self *SAzureClient) PerformAction(resourceId string, action string, body s
return nil, err
}
url := fmt.Sprintf("%s/%s", resourceId, action)
return jsonRequest(cli, "POST", self.domain, url, body)
return jsonRequest(cli, "POST", self.domain, url, self.subscriptionId, body)
}
func (self *SAzureClient) fetchResourceGroup(cli *autorest.Client, location string) error {
@@ -261,7 +263,7 @@ func (self *SAzureClient) fetchResourceGroup(cli *autorest.Client, location stri
if len(self.ressourceGroups) == 0 {
//Create Default resourceGroup
_url := fmt.Sprintf("/subscriptions/%s/resourcegroups/Default", self.subscriptionId)
body, err := jsonRequest(cli, "PUT", self.domain, _url, fmt.Sprintf(`{"name": "Default", "location": "%s"}`, location))
body, err := jsonRequest(cli, "PUT", self.domain, _url, self.subscriptionId, fmt.Sprintf(`{"name": "Default", "location": "%s"}`, location))
if err != nil {
return err
}
@@ -287,6 +289,18 @@ func (self *SAzureClient) checkParams(body jsonutils.JSONObject, params []string
return result, nil
}
type AzureErrorDetail struct {
Code string `json:"code,omitempty"`
Message string `json:"message,omitempty"`
Target string `json:"target,omitempty"`
}
type AzureError struct {
Code string `json:"code,omitempty"`
Details []AzureErrorDetail `json:"details,omitempty"`
Message string `json:"message,omitempty"`
}
func (self *SAzureClient) Create(body jsonutils.JSONObject, retVal interface{}) error {
cli, err := self.getDefaultClient()
if err != nil {
@@ -307,11 +321,14 @@ func (self *SAzureClient) Create(body jsonutils.JSONObject, retVal interface{})
return fmt.Errorf("Create Default resourceGroup error?")
}
url := fmt.Sprintf("/subscriptions/%s/resourceGroups/%s/providers/%s/%s", self.subscriptionId, self.ressourceGroups[0].Name, params["type"], params["name"])
result, err := jsonRequest(cli, "PUT", self.domain, url, body.String())
result, err := jsonRequest(cli, "PUT", self.domain, url, self.subscriptionId, body.String())
if err != nil {
return err
}
return result.Unmarshal(retVal)
if retVal != nil {
return result.Unmarshal(retVal)
}
return nil
}
func (self *SAzureClient) CheckNameAvailability(Type string, body string) (jsonutils.JSONObject, error) {
@@ -323,7 +340,7 @@ func (self *SAzureClient) CheckNameAvailability(Type string, body string) (jsonu
return nil, fmt.Errorf("Missing subscription ID")
}
url := fmt.Sprintf("/subscriptions/%s/providers/%s/checkNameAvailability", self.subscriptionId, Type)
return jsonRequest(cli, "POST", self.domain, url, body)
return jsonRequest(cli, "POST", self.domain, url, self.subscriptionId, body)
}
func (self *SAzureClient) Update(body jsonutils.JSONObject, retVal interface{}) error {
@@ -332,7 +349,7 @@ func (self *SAzureClient) Update(body jsonutils.JSONObject, retVal interface{})
return err
}
url, err := body.GetString("id")
result, err := jsonRequest(cli, "PUT", self.domain, url, body.String())
result, err := jsonRequest(cli, "PUT", self.domain, url, self.subscriptionId, body.String())
if err != nil {
return err
}
@@ -342,8 +359,86 @@ func (self *SAzureClient) Update(body jsonutils.JSONObject, retVal interface{})
return nil
}
func jsonRequest(client *autorest.Client, method, domain, baseUrl string, body string) (jsonutils.JSONObject, error) {
return _jsonRequest(client, method, domain, baseUrl, body)
func waitRegisterComplete(client *autorest.Client, domain, subscriptionId string, serviceType string) error {
for i := 1; i < 10; i++ {
result, err := _jsonRequest(client, "GET", domain, fmt.Sprintf("/subscriptions/%s/providers", subscriptionId), "")
if err != nil {
return err
}
value, err := result.GetArray("value")
if err != nil {
return err
}
for _, v := range value {
namespace, _ := v.GetString("namespace")
if namespace == serviceType {
state, _ := v.GetString("registrationState")
if state == "Registered" {
return nil
}
log.Debugf("service %s state %s waite %d second ...", serviceType, state, i*10)
}
}
time.Sleep(time.Second * time.Duration(i*10))
}
return fmt.Errorf("wait service %s register timeout", serviceType)
}
func registerService(client *autorest.Client, domain, subscriptionId string, serviceType string) error {
registryUrl := fmt.Sprintf("/subscriptions/%s/providers/%s/register", subscriptionId, serviceType)
result, err := _jsonRequest(client, "POST", domain, registryUrl, "")
if err != nil || result.Contains("error") {
return fmt.Errorf("failed to register %s service", serviceType)
}
if state, _ := result.GetString("registrationState"); state == "Registered" {
return nil
}
return waitRegisterComplete(client, domain, subscriptionId, serviceType)
}
func recoverFromError(client *autorest.Client, domain, subscriptionId string, azureErr AzureError) bool {
switch azureErr.Code {
case "SubscriptionNotRegistered":
services := []string{"Microsoft.Network"}
for _, service := range services {
if err := registerService(client, domain, subscriptionId, service); err != nil {
log.Errorf("register %s error: %v", service, err)
return false
}
}
return true
case "MissingSubscriptionRegistration":
for _, detail := range azureErr.Details {
log.Errorf("The subscription is not registered to use namespace '%s', try register it", detail.Target)
if err := registerService(client, domain, subscriptionId, detail.Target); err != nil {
log.Errorf("register %s error: %v", detail.Target, err)
return false
}
}
return true
default:
return false
}
return false
}
func jsonRequest(client *autorest.Client, method, domain, baseUrl string, subscriptionId string, body string) (jsonutils.JSONObject, error) {
result, err := _jsonRequest(client, method, domain, baseUrl, body)
if err != nil {
return nil, err
}
if result.Contains("error") {
azureError := AzureError{}
if err := result.Unmarshal(&azureError, "error"); err != nil {
return nil, fmt.Errorf(result.String())
}
if recoverFromError(client, domain, subscriptionId, azureError) {
return _jsonRequest(client, method, domain, baseUrl, body)
}
log.Errorf("Azure %s request: %s \nbody: %s error: %v", method, baseUrl, body, result.String())
return nil, fmt.Errorf(result.String())
}
return result, nil
}
func waitForComplatetion(client *autorest.Client, req *http.Request, resp *http.Response, timeout time.Duration) (jsonutils.JSONObject, error) {
@@ -470,15 +565,7 @@ func _jsonRequest(client *autorest.Client, method, domain, baseURL, body string)
return nil, err
}
_data := strings.Replace(string(data), "\r", "", -1)
result, err = jsonutils.Parse([]byte(_data))
if err != nil {
return nil, err
}
if result.Contains("error") {
log.Errorf("Azure %s request: %s \nbody: %s error: %v", req.Method, req.URL.String(), body, result.String())
return nil, fmt.Errorf(result.String())
}
return result, nil
return jsonutils.Parse([]byte(_data))
}
func (self *SAzureClient) UpdateAccount(tenantId, secret, envName string) error {
+1 -5
View File
@@ -13,10 +13,6 @@ type SClassicHost struct {
zone *SZone
}
const (
DEFAULT_USER = "yunion"
)
func (self *SClassicHost) GetMetadata() *jsonutils.JSONDict {
return nil
}
@@ -26,7 +22,7 @@ func (self *SClassicHost) GetId() string {
}
func (self *SClassicHost) GetName() string {
return fmt.Sprintf("%s-classic", self.zone.region.client.subscriptionId)
return fmt.Sprintf("%s/%s-classic", self.zone.region.GetGlobalId(), self.zone.region.client.subscriptionId)
}
func (self *SClassicHost) GetGlobalId() string {
+2 -1
View File
@@ -114,7 +114,8 @@ type SClassicInstance struct {
func (self *SClassicInstance) GetMetadata() *jsonutils.JSONDict {
data := jsonutils.NewDict()
data.Add(jsonutils.NewString(self.Properties.HardwareProfile.Size), "price_key")
priceKey := fmt.Sprintf("%s::%s", self.Properties.HardwareProfile.Size, self.host.zone.region.Name)
data.Add(jsonutils.NewString(priceKey), "price_key")
if self.Properties.NetworkProfile.NetworkSecurityGroup != nil {
data.Add(jsonutils.NewString(self.Properties.NetworkProfile.NetworkSecurityGroup.ID), "secgroupId")
}
+4
View File
@@ -9,6 +9,10 @@ import (
"github.com/Azure/go-autorest/autorest"
)
const (
DEBUG = false
)
func LogRequest() autorest.PrepareDecorator {
return func(p autorest.Preparer) autorest.Preparer {
return autorest.PreparerFunc(func(r *http.Request) (*http.Request, error) {
+43 -13
View File
@@ -8,6 +8,7 @@ import (
"yunion.io/x/log"
"yunion.io/x/onecloud/pkg/cloudprovider"
"yunion.io/x/onecloud/pkg/compute/models"
"yunion.io/x/onecloud/pkg/util/ansible"
)
type SHost struct {
@@ -23,7 +24,7 @@ func (self *SHost) GetId() string {
}
func (self *SHost) GetName() string {
return fmt.Sprintf("%s", self.zone.region.client.subscriptionId)
return fmt.Sprintf("%s/%s", self.zone.region.GetGlobalId(), self.zone.region.client.subscriptionId)
}
func (self *SHost) GetGlobalId() string {
@@ -42,18 +43,47 @@ func (self *SHost) Refresh() error {
return nil
}
func (self *SHost) CreateVM(name string, imgId string, sysDiskSize int, cpu int, memMB int, networkId string, ipAddr string, desc string, passwd string, storageType string, diskSizes []int, publicKey string, secgroupId string, userData string) (cloudprovider.ICloudVM, error) {
nicId := ""
if net := self.zone.getNetworkById(networkId); net == nil {
return nil, fmt.Errorf("invalid network ID %s", networkId)
} else if nic, err := self.zone.region.CreateNetworkInterface(fmt.Sprintf("%s-ipconfig", name), ipAddr, net.GetId(), secgroupId); err != nil {
return nil, err
} else {
nicId = nic.ID
}
vmId, err := self._createVM(name, imgId, int32(sysDiskSize), cpu, memMB, nicId, ipAddr, desc, passwd, storageType, diskSizes, publicKey, userData)
func (self *SHost) searchNetorkInterface(IPAddr string, networkId string, secgroupId string) (*SInstanceNic, error) {
interfaces, err := self.zone.region.GetNetworkInterfaces()
if err != nil {
self.zone.region.DeleteNetworkInterface(nicId)
return nil, err
}
for i, nic := range interfaces {
for _, ipConf := range nic.Properties.IPConfigurations {
if ipConf.Properties.PrivateIPAddress == IPAddr && networkId == ipConf.Properties.Subnet.ID && ipConf.Properties.PrivateIPAllocationMethod == "Static" {
if nic.Properties.NetworkSecurityGroup == nil || nic.Properties.NetworkSecurityGroup.ID != secgroupId {
nic.Properties.NetworkSecurityGroup = &SSecurityGroup{ID: secgroupId}
if err := self.zone.region.client.Update(jsonutils.Marshal(nic), nil); err != nil {
log.Errorf("assign secgroup %s for nic %s failed %d")
return nil, err
}
}
return &interfaces[i], nil
}
}
}
return nil, cloudprovider.ErrNotFound
}
func (self *SHost) CreateVM(name string, imgId string, sysDiskSize int, cpu int, memMB int, networkId string, ipAddr string, desc string, passwd string, storageType string, diskSizes []int, publicKey string, secgroupId string, userData string) (cloudprovider.ICloudVM, error) {
net := self.zone.getNetworkById(networkId)
if net == nil {
return nil, fmt.Errorf("invalid network ID %s", networkId)
}
nic, err := self.searchNetorkInterface(ipAddr, net.GetId(), secgroupId)
if err != nil {
if err == cloudprovider.ErrNotFound {
nic, err = self.zone.region.CreateNetworkInterface(fmt.Sprintf("%s-ipconfig", name), ipAddr, net.GetId(), secgroupId)
if err != nil {
return nil, err
}
} else {
return nil, err
}
}
vmId, err := self._createVM(name, imgId, int32(sysDiskSize), cpu, memMB, nic.ID, ipAddr, desc, passwd, storageType, diskSizes, publicKey, userData)
if err != nil {
self.zone.region.DeleteNetworkInterface(nic.ID)
return nil, err
}
if vm, err := self.zone.region.GetInstance(vmId); err != nil {
@@ -89,7 +119,7 @@ func (self *SHost) _createVM(name string, imgId string, sysDiskSize int32, cpu i
},
OsProfile: OsProfile{
ComputerName: name,
AdminUsername: DEFAULT_USER,
AdminUsername: ansible.PUBLIC_CLOUD_ANSIBLE_USER,
AdminPassword: passwd,
CustomData: userData,
},
+2 -1
View File
@@ -230,7 +230,8 @@ func (self *SInstance) GetMetadata() *jsonutils.JSONDict {
data.Add(jsonutils.NewString(loginKey), "login_key")
}
data.Add(jsonutils.NewString(self.Properties.HardwareProfile.VMSize), "price_key")
priceKey := fmt.Sprintf("%s::%s", self.Properties.HardwareProfile.VMSize, self.host.zone.region.Name)
data.Add(jsonutils.NewString(priceKey), "price_key")
if nics, err := self.getNics(); err == nil {
for _, nic := range nics {
if nic.Properties.NetworkSecurityGroup != nil {
+2 -1
View File
@@ -2,6 +2,7 @@ package azure
import (
"fmt"
"yunion.io/x/pkg/utils"
)
@@ -108,7 +109,7 @@ func (self *SAzureClient) ListResourceSkus() ([]SResourceSku, error) {
url := fmt.Sprintf("/subscriptions/%s/providers/Microsoft.Compute/skus?api-version=2017-09-01", self.subscriptionId)
skus := make([]SResourceSku, 0)
for {
body, err := jsonRequest(cli, "GET", self.domain, url, "")
body, err := jsonRequest(cli, "GET", self.domain, url, self.subscriptionId, "")
if err != nil {
return nil, err
}
+55
View File
@@ -0,0 +1,55 @@
package azure
import (
"fmt"
)
type SServices struct {
Value []SService `json:"value,omitempty"`
}
type SService struct {
ID string `json:"id,omitempty"`
Namespace string `json:"namespace,omitempty"`
RegistrationState string `json:"registrationState,omitempty"`
ResourceTypes []ResourceType `json:"resourceTypes,omitempty"`
}
type ResourceType struct {
ApiVersions []string `json:"apiVersions,omitempty"`
Capabilities string `json:"capabilities,omitempty"`
Locations []string `json:"locations,omitempty"`
ResourceType string `json:"locations,omitempty"`
}
func (self *SRegion) ListServices() ([]SService, error) {
services := []SService{}
return services, self.client.List("providers", &services)
}
func (self *SRegion) SerciceShow(serviceType string) (*SService, error) {
service := SService{}
return &service, self.client.Get("providers/"+serviceType, []string{}, &service)
}
func (self *SRegion) serviceOperation(resourceType, operation string) error {
services, err := self.ListServices()
if err != nil {
return err
}
for _, service := range services {
if service.Namespace == resourceType {
_, err := self.client.jsonRequest("POST", fmt.Sprintf("%s/%s", service.ID, operation), "")
return err
}
}
return fmt.Errorf("failed to find namespace: %s", resourceType)
}
func (self *SRegion) ServiceRegister(resourceType string) error {
return self.serviceOperation(resourceType, "register")
}
func (self *SRegion) ServiceUnRegister(resourceType string) error {
return self.serviceOperation(resourceType, "unregister")
}
+41
View File
@@ -0,0 +1,41 @@
package shell
import (
"yunion.io/x/onecloud/pkg/util/azure"
"yunion.io/x/onecloud/pkg/util/shellutils"
)
func init() {
type ServiceListOptions struct {
}
shellutils.R(&ServiceListOptions{}, "service-list", "List providers", func(cli *azure.SRegion, args *ServiceListOptions) error {
services, err := cli.ListServices()
if err != nil {
return err
}
printList(services, len(services), 0, 0, []string{})
return nil
})
type ServiceOptions struct {
NAME string `help:"Name for service register"`
}
shellutils.R(&ServiceOptions{}, "service-register", "Register service", func(cli *azure.SRegion, args *ServiceOptions) error {
return cli.ServiceRegister(args.NAME)
})
shellutils.R(&ServiceOptions{}, "service-unregister", "Unregister service", func(cli *azure.SRegion, args *ServiceOptions) error {
return cli.ServiceUnRegister(args.NAME)
})
shellutils.R(&ServiceOptions{}, "service-show", "Show service detail", func(cli *azure.SRegion, args *ServiceOptions) error {
service, err := cli.SerciceShow(args.NAME)
if err != nil {
return err
}
printObject(service)
return nil
})
}
+19
View File
@@ -32,4 +32,23 @@ func init() {
printObject(vpc)
return nil
})
shellutils.R(&VpcOptions{}, "vpc-delete", "Delete vpc", func(cli *azure.SRegion, args *VpcOptions) error {
return cli.DeleteVpc(args.ID)
})
type VpcCreateOptions struct {
NAME string `help:"vpc Name"`
CIDR string `help:"vpc cidr"`
Desc string `help:"vpc description"`
}
shellutils.R(&VpcCreateOptions{}, "vpc-create", "Create vpc", func(cli *azure.SRegion, args *VpcCreateOptions) error {
vpc, err := cli.CreateIVpc(args.NAME, args.Desc, args.CIDR)
if err != nil {
return err
}
printObject(vpc)
return nil
})
}
+2 -8
View File
@@ -154,9 +154,11 @@ func (self *SRegion) GetISnapshots() ([]cloudprovider.ICloudSnapshot, error) {
classicSnapshots = append(classicSnapshots, _classicSnapshots...)
isnapshots := make([]cloudprovider.ICloudSnapshot, len(snapshots)+len(classicSnapshots))
for i := 0; i < len(snapshots); i++ {
snapshots[i].region = self
isnapshots[i] = &snapshots[i]
}
for i := 0; i < len(classicSnapshots); i++ {
classicSnapshots[i].region = self
isnapshots[len(snapshots)+i] = &classicSnapshots[i]
}
return isnapshots, nil
@@ -166,14 +168,6 @@ func (self *SSnapshot) GetDiskId() string {
return self.Properties.CreationData.SourceResourceID
}
func (self *SSnapshot) GetManagerId() string {
return self.region.client.providerId
}
func (self *SSnapshot) GetRegionId() string {
return self.region.GetId()
}
func (self *SSnapshot) GetDiskType() string {
return ""
}
+5 -1
View File
@@ -77,7 +77,11 @@ func (self *SVpc) GetCidrBlock() string {
}
func (self *SVpc) Delete() error {
return self.region.client.Delete(self.ID)
return self.region.DeleteVpc(self.ID)
}
func (self *SRegion) DeleteVpc(vpcId string) error {
return self.client.Delete(vpcId)
}
func (self *SVpc) getSecurityGroups() ([]SSecurityGroup, error) {
+33 -23
View File
@@ -13,29 +13,29 @@ type TRbacResult string
const (
WILD_MATCH = "*"
Allow = TRbacResult("allow")
AdminAllow = TRbacResult("allow")
OwnerAllow = TRbacResult("owner")
UserAllow = TRbacResult("user")
GuestAllow = TRbacResult("guest")
Deny = TRbacResult("deny")
)
func (result TRbacResult) IsHigherPrivilege(r2 TRbacResult) bool {
switch result {
case Allow:
if r2 == Allow {
return false
} else {
return true
}
case OwnerAllow:
if r2 == Deny {
return true
} else {
return false
}
case Deny:
return false
var (
strictness = map[TRbacResult]int{
Deny: 0,
AdminAllow: 1,
OwnerAllow: 2,
UserAllow: 3,
GuestAllow: 4,
}
return false
)
func (r TRbacResult) Strictness() int {
return strictness[r]
}
func (r1 TRbacResult) StricterThan(r2 TRbacResult) bool {
return r1.Strictness() < r2.Strictness()
}
type SRbacPolicy struct {
@@ -79,6 +79,10 @@ func (rule *SRbacRule) contains(rule2 *SRbacRule) bool {
return true
}
func (rule *SRbacRule) stricterThan(r2 *SRbacRule) bool {
return rule.Result.StricterThan(r2.Result)
}
func (rule *SRbacRule) match(service string, resource string, action string, extra ...string) (bool, int, int) {
matched := 0
weight := 0
@@ -144,7 +148,9 @@ func (policy *SRbacPolicy) GetMatchRule(service string, resource string, action
var matchRule *SRbacRule
for i := 0; i < len(policy.Rules); i += 1 {
match, matchCnt, weight := policy.Rules[i].match(service, resource, action, extra...)
if match && (maxMatchCnt < matchCnt || (maxMatchCnt == matchCnt && minWeight > weight)) {
if match && (maxMatchCnt < matchCnt ||
(maxMatchCnt == matchCnt && minWeight > weight) ||
(maxMatchCnt == matchCnt && minWeight == weight && matchRule.stricterThan(&policy.Rules[i]))) {
maxMatchCnt = matchCnt
minWeight = weight
matchRule = &policy.Rules[i]
@@ -153,7 +159,7 @@ func (policy *SRbacPolicy) GetMatchRule(service string, resource string, action
return matchRule
}
func compactRules(rules []SRbacRule) []SRbacRule {
func CompactRules(rules []SRbacRule) []SRbacRule {
output := make([]SRbacRule, 1)
output[0] = rules[0]
for i := 1; i < len(rules); i += 1 {
@@ -190,7 +196,7 @@ func (policy *SRbacPolicy) Decode(policyJson jsonutils.JSONObject) error {
return err
}
policy.Rules = compactRules(rules)
policy.Rules = CompactRules(rules)
return nil
}
@@ -208,10 +214,14 @@ func decode(rules jsonutils.JSONObject, decodeRule SRbacRule, level int) ([]SRba
ruleJsonStr := rules.(*jsonutils.JSONString)
ruleStr, _ := ruleJsonStr.GetString()
switch ruleStr {
case string(Allow):
decodeRule.Result = Allow
case string(AdminAllow):
decodeRule.Result = AdminAllow
case string(OwnerAllow):
decodeRule.Result = OwnerAllow
case string(UserAllow):
decodeRule.Result = UserAllow
case string(GuestAllow):
decodeRule.Result = GuestAllow
case string(Deny):
decodeRule.Result = Deny
default:
+1 -1
View File
@@ -31,7 +31,7 @@ type SSHtoolSol struct {
func getCommand(ctx context.Context, userCred mcclient.TokenCredential, ip string) (string, *BaseCommand, error) {
cmd := NewBaseCommand(o.Options.SshToolPath)
s := auth.GetAdminSession(o.Options.Region, "v2")
key, err := modules.Sshkeypairs.GetById(s, userCred.GetProjectId(), jsonutils.NewDict())
key, err := modules.Sshkeypairs.GetById(s, userCred.GetProjectId(), jsonutils.Marshal(map[string]bool{"admin": true}))
if err != nil {
return "", nil, err
}
+403
View File
@@ -0,0 +1,403 @@
// Code generated by private/model/cli/gen-api/main.go. DO NOT EDIT.
// Package s3iface provides an interface to enable mocking the Amazon Simple Storage Service service client
// for testing your code.
//
// It is important to note that this interface will have breaking changes
// when the service model is updated and adds new API operations, paginators,
// and waiters.
package s3iface
import (
"github.com/aws/aws-sdk-go/aws"
"github.com/aws/aws-sdk-go/aws/request"
"github.com/aws/aws-sdk-go/service/s3"
)
// S3API provides an interface to enable mocking the
// s3.S3 service client's API operation,
// paginators, and waiters. This make unit testing your code that calls out
// to the SDK's service client's calls easier.
//
// The best way to use this interface is so the SDK's service client's calls
// can be stubbed out for unit testing your code with the SDK without needing
// to inject custom request handlers into the SDK's request pipeline.
//
// // myFunc uses an SDK service client to make a request to
// // Amazon Simple Storage Service.
// func myFunc(svc s3iface.S3API) bool {
// // Make svc.AbortMultipartUpload request
// }
//
// func main() {
// sess := session.New()
// svc := s3.New(sess)
//
// myFunc(svc)
// }
//
// In your _test.go file:
//
// // Define a mock struct to be used in your unit tests of myFunc.
// type mockS3Client struct {
// s3iface.S3API
// }
// func (m *mockS3Client) AbortMultipartUpload(input *s3.AbortMultipartUploadInput) (*s3.AbortMultipartUploadOutput, error) {
// // mock response/functionality
// }
//
// func TestMyFunc(t *testing.T) {
// // Setup Test
// mockSvc := &mockS3Client{}
//
// myfunc(mockSvc)
//
// // Verify myFunc's functionality
// }
//
// It is important to note that this interface will have breaking changes
// when the service model is updated and adds new API operations, paginators,
// and waiters. Its suggested to use the pattern above for testing, or using
// tooling to generate mocks to satisfy the interfaces.
type S3API interface {
AbortMultipartUpload(*s3.AbortMultipartUploadInput) (*s3.AbortMultipartUploadOutput, error)
AbortMultipartUploadWithContext(aws.Context, *s3.AbortMultipartUploadInput, ...request.Option) (*s3.AbortMultipartUploadOutput, error)
AbortMultipartUploadRequest(*s3.AbortMultipartUploadInput) (*request.Request, *s3.AbortMultipartUploadOutput)
CompleteMultipartUpload(*s3.CompleteMultipartUploadInput) (*s3.CompleteMultipartUploadOutput, error)
CompleteMultipartUploadWithContext(aws.Context, *s3.CompleteMultipartUploadInput, ...request.Option) (*s3.CompleteMultipartUploadOutput, error)
CompleteMultipartUploadRequest(*s3.CompleteMultipartUploadInput) (*request.Request, *s3.CompleteMultipartUploadOutput)
CopyObject(*s3.CopyObjectInput) (*s3.CopyObjectOutput, error)
CopyObjectWithContext(aws.Context, *s3.CopyObjectInput, ...request.Option) (*s3.CopyObjectOutput, error)
CopyObjectRequest(*s3.CopyObjectInput) (*request.Request, *s3.CopyObjectOutput)
CreateBucket(*s3.CreateBucketInput) (*s3.CreateBucketOutput, error)
CreateBucketWithContext(aws.Context, *s3.CreateBucketInput, ...request.Option) (*s3.CreateBucketOutput, error)
CreateBucketRequest(*s3.CreateBucketInput) (*request.Request, *s3.CreateBucketOutput)
CreateMultipartUpload(*s3.CreateMultipartUploadInput) (*s3.CreateMultipartUploadOutput, error)
CreateMultipartUploadWithContext(aws.Context, *s3.CreateMultipartUploadInput, ...request.Option) (*s3.CreateMultipartUploadOutput, error)
CreateMultipartUploadRequest(*s3.CreateMultipartUploadInput) (*request.Request, *s3.CreateMultipartUploadOutput)
DeleteBucket(*s3.DeleteBucketInput) (*s3.DeleteBucketOutput, error)
DeleteBucketWithContext(aws.Context, *s3.DeleteBucketInput, ...request.Option) (*s3.DeleteBucketOutput, error)
DeleteBucketRequest(*s3.DeleteBucketInput) (*request.Request, *s3.DeleteBucketOutput)
DeleteBucketAnalyticsConfiguration(*s3.DeleteBucketAnalyticsConfigurationInput) (*s3.DeleteBucketAnalyticsConfigurationOutput, error)
DeleteBucketAnalyticsConfigurationWithContext(aws.Context, *s3.DeleteBucketAnalyticsConfigurationInput, ...request.Option) (*s3.DeleteBucketAnalyticsConfigurationOutput, error)
DeleteBucketAnalyticsConfigurationRequest(*s3.DeleteBucketAnalyticsConfigurationInput) (*request.Request, *s3.DeleteBucketAnalyticsConfigurationOutput)
DeleteBucketCors(*s3.DeleteBucketCorsInput) (*s3.DeleteBucketCorsOutput, error)
DeleteBucketCorsWithContext(aws.Context, *s3.DeleteBucketCorsInput, ...request.Option) (*s3.DeleteBucketCorsOutput, error)
DeleteBucketCorsRequest(*s3.DeleteBucketCorsInput) (*request.Request, *s3.DeleteBucketCorsOutput)
DeleteBucketEncryption(*s3.DeleteBucketEncryptionInput) (*s3.DeleteBucketEncryptionOutput, error)
DeleteBucketEncryptionWithContext(aws.Context, *s3.DeleteBucketEncryptionInput, ...request.Option) (*s3.DeleteBucketEncryptionOutput, error)
DeleteBucketEncryptionRequest(*s3.DeleteBucketEncryptionInput) (*request.Request, *s3.DeleteBucketEncryptionOutput)
DeleteBucketInventoryConfiguration(*s3.DeleteBucketInventoryConfigurationInput) (*s3.DeleteBucketInventoryConfigurationOutput, error)
DeleteBucketInventoryConfigurationWithContext(aws.Context, *s3.DeleteBucketInventoryConfigurationInput, ...request.Option) (*s3.DeleteBucketInventoryConfigurationOutput, error)
DeleteBucketInventoryConfigurationRequest(*s3.DeleteBucketInventoryConfigurationInput) (*request.Request, *s3.DeleteBucketInventoryConfigurationOutput)
DeleteBucketLifecycle(*s3.DeleteBucketLifecycleInput) (*s3.DeleteBucketLifecycleOutput, error)
DeleteBucketLifecycleWithContext(aws.Context, *s3.DeleteBucketLifecycleInput, ...request.Option) (*s3.DeleteBucketLifecycleOutput, error)
DeleteBucketLifecycleRequest(*s3.DeleteBucketLifecycleInput) (*request.Request, *s3.DeleteBucketLifecycleOutput)
DeleteBucketMetricsConfiguration(*s3.DeleteBucketMetricsConfigurationInput) (*s3.DeleteBucketMetricsConfigurationOutput, error)
DeleteBucketMetricsConfigurationWithContext(aws.Context, *s3.DeleteBucketMetricsConfigurationInput, ...request.Option) (*s3.DeleteBucketMetricsConfigurationOutput, error)
DeleteBucketMetricsConfigurationRequest(*s3.DeleteBucketMetricsConfigurationInput) (*request.Request, *s3.DeleteBucketMetricsConfigurationOutput)
DeleteBucketPolicy(*s3.DeleteBucketPolicyInput) (*s3.DeleteBucketPolicyOutput, error)
DeleteBucketPolicyWithContext(aws.Context, *s3.DeleteBucketPolicyInput, ...request.Option) (*s3.DeleteBucketPolicyOutput, error)
DeleteBucketPolicyRequest(*s3.DeleteBucketPolicyInput) (*request.Request, *s3.DeleteBucketPolicyOutput)
DeleteBucketReplication(*s3.DeleteBucketReplicationInput) (*s3.DeleteBucketReplicationOutput, error)
DeleteBucketReplicationWithContext(aws.Context, *s3.DeleteBucketReplicationInput, ...request.Option) (*s3.DeleteBucketReplicationOutput, error)
DeleteBucketReplicationRequest(*s3.DeleteBucketReplicationInput) (*request.Request, *s3.DeleteBucketReplicationOutput)
DeleteBucketTagging(*s3.DeleteBucketTaggingInput) (*s3.DeleteBucketTaggingOutput, error)
DeleteBucketTaggingWithContext(aws.Context, *s3.DeleteBucketTaggingInput, ...request.Option) (*s3.DeleteBucketTaggingOutput, error)
DeleteBucketTaggingRequest(*s3.DeleteBucketTaggingInput) (*request.Request, *s3.DeleteBucketTaggingOutput)
DeleteBucketWebsite(*s3.DeleteBucketWebsiteInput) (*s3.DeleteBucketWebsiteOutput, error)
DeleteBucketWebsiteWithContext(aws.Context, *s3.DeleteBucketWebsiteInput, ...request.Option) (*s3.DeleteBucketWebsiteOutput, error)
DeleteBucketWebsiteRequest(*s3.DeleteBucketWebsiteInput) (*request.Request, *s3.DeleteBucketWebsiteOutput)
DeleteObject(*s3.DeleteObjectInput) (*s3.DeleteObjectOutput, error)
DeleteObjectWithContext(aws.Context, *s3.DeleteObjectInput, ...request.Option) (*s3.DeleteObjectOutput, error)
DeleteObjectRequest(*s3.DeleteObjectInput) (*request.Request, *s3.DeleteObjectOutput)
DeleteObjectTagging(*s3.DeleteObjectTaggingInput) (*s3.DeleteObjectTaggingOutput, error)
DeleteObjectTaggingWithContext(aws.Context, *s3.DeleteObjectTaggingInput, ...request.Option) (*s3.DeleteObjectTaggingOutput, error)
DeleteObjectTaggingRequest(*s3.DeleteObjectTaggingInput) (*request.Request, *s3.DeleteObjectTaggingOutput)
DeleteObjects(*s3.DeleteObjectsInput) (*s3.DeleteObjectsOutput, error)
DeleteObjectsWithContext(aws.Context, *s3.DeleteObjectsInput, ...request.Option) (*s3.DeleteObjectsOutput, error)
DeleteObjectsRequest(*s3.DeleteObjectsInput) (*request.Request, *s3.DeleteObjectsOutput)
GetBucketAccelerateConfiguration(*s3.GetBucketAccelerateConfigurationInput) (*s3.GetBucketAccelerateConfigurationOutput, error)
GetBucketAccelerateConfigurationWithContext(aws.Context, *s3.GetBucketAccelerateConfigurationInput, ...request.Option) (*s3.GetBucketAccelerateConfigurationOutput, error)
GetBucketAccelerateConfigurationRequest(*s3.GetBucketAccelerateConfigurationInput) (*request.Request, *s3.GetBucketAccelerateConfigurationOutput)
GetBucketAcl(*s3.GetBucketAclInput) (*s3.GetBucketAclOutput, error)
GetBucketAclWithContext(aws.Context, *s3.GetBucketAclInput, ...request.Option) (*s3.GetBucketAclOutput, error)
GetBucketAclRequest(*s3.GetBucketAclInput) (*request.Request, *s3.GetBucketAclOutput)
GetBucketAnalyticsConfiguration(*s3.GetBucketAnalyticsConfigurationInput) (*s3.GetBucketAnalyticsConfigurationOutput, error)
GetBucketAnalyticsConfigurationWithContext(aws.Context, *s3.GetBucketAnalyticsConfigurationInput, ...request.Option) (*s3.GetBucketAnalyticsConfigurationOutput, error)
GetBucketAnalyticsConfigurationRequest(*s3.GetBucketAnalyticsConfigurationInput) (*request.Request, *s3.GetBucketAnalyticsConfigurationOutput)
GetBucketCors(*s3.GetBucketCorsInput) (*s3.GetBucketCorsOutput, error)
GetBucketCorsWithContext(aws.Context, *s3.GetBucketCorsInput, ...request.Option) (*s3.GetBucketCorsOutput, error)
GetBucketCorsRequest(*s3.GetBucketCorsInput) (*request.Request, *s3.GetBucketCorsOutput)
GetBucketEncryption(*s3.GetBucketEncryptionInput) (*s3.GetBucketEncryptionOutput, error)
GetBucketEncryptionWithContext(aws.Context, *s3.GetBucketEncryptionInput, ...request.Option) (*s3.GetBucketEncryptionOutput, error)
GetBucketEncryptionRequest(*s3.GetBucketEncryptionInput) (*request.Request, *s3.GetBucketEncryptionOutput)
GetBucketInventoryConfiguration(*s3.GetBucketInventoryConfigurationInput) (*s3.GetBucketInventoryConfigurationOutput, error)
GetBucketInventoryConfigurationWithContext(aws.Context, *s3.GetBucketInventoryConfigurationInput, ...request.Option) (*s3.GetBucketInventoryConfigurationOutput, error)
GetBucketInventoryConfigurationRequest(*s3.GetBucketInventoryConfigurationInput) (*request.Request, *s3.GetBucketInventoryConfigurationOutput)
GetBucketLifecycle(*s3.GetBucketLifecycleInput) (*s3.GetBucketLifecycleOutput, error)
GetBucketLifecycleWithContext(aws.Context, *s3.GetBucketLifecycleInput, ...request.Option) (*s3.GetBucketLifecycleOutput, error)
GetBucketLifecycleRequest(*s3.GetBucketLifecycleInput) (*request.Request, *s3.GetBucketLifecycleOutput)
GetBucketLifecycleConfiguration(*s3.GetBucketLifecycleConfigurationInput) (*s3.GetBucketLifecycleConfigurationOutput, error)
GetBucketLifecycleConfigurationWithContext(aws.Context, *s3.GetBucketLifecycleConfigurationInput, ...request.Option) (*s3.GetBucketLifecycleConfigurationOutput, error)
GetBucketLifecycleConfigurationRequest(*s3.GetBucketLifecycleConfigurationInput) (*request.Request, *s3.GetBucketLifecycleConfigurationOutput)
GetBucketLocation(*s3.GetBucketLocationInput) (*s3.GetBucketLocationOutput, error)
GetBucketLocationWithContext(aws.Context, *s3.GetBucketLocationInput, ...request.Option) (*s3.GetBucketLocationOutput, error)
GetBucketLocationRequest(*s3.GetBucketLocationInput) (*request.Request, *s3.GetBucketLocationOutput)
GetBucketLogging(*s3.GetBucketLoggingInput) (*s3.GetBucketLoggingOutput, error)
GetBucketLoggingWithContext(aws.Context, *s3.GetBucketLoggingInput, ...request.Option) (*s3.GetBucketLoggingOutput, error)
GetBucketLoggingRequest(*s3.GetBucketLoggingInput) (*request.Request, *s3.GetBucketLoggingOutput)
GetBucketMetricsConfiguration(*s3.GetBucketMetricsConfigurationInput) (*s3.GetBucketMetricsConfigurationOutput, error)
GetBucketMetricsConfigurationWithContext(aws.Context, *s3.GetBucketMetricsConfigurationInput, ...request.Option) (*s3.GetBucketMetricsConfigurationOutput, error)
GetBucketMetricsConfigurationRequest(*s3.GetBucketMetricsConfigurationInput) (*request.Request, *s3.GetBucketMetricsConfigurationOutput)
GetBucketNotification(*s3.GetBucketNotificationConfigurationRequest) (*s3.NotificationConfigurationDeprecated, error)
GetBucketNotificationWithContext(aws.Context, *s3.GetBucketNotificationConfigurationRequest, ...request.Option) (*s3.NotificationConfigurationDeprecated, error)
GetBucketNotificationRequest(*s3.GetBucketNotificationConfigurationRequest) (*request.Request, *s3.NotificationConfigurationDeprecated)
GetBucketNotificationConfiguration(*s3.GetBucketNotificationConfigurationRequest) (*s3.NotificationConfiguration, error)
GetBucketNotificationConfigurationWithContext(aws.Context, *s3.GetBucketNotificationConfigurationRequest, ...request.Option) (*s3.NotificationConfiguration, error)
GetBucketNotificationConfigurationRequest(*s3.GetBucketNotificationConfigurationRequest) (*request.Request, *s3.NotificationConfiguration)
GetBucketPolicy(*s3.GetBucketPolicyInput) (*s3.GetBucketPolicyOutput, error)
GetBucketPolicyWithContext(aws.Context, *s3.GetBucketPolicyInput, ...request.Option) (*s3.GetBucketPolicyOutput, error)
GetBucketPolicyRequest(*s3.GetBucketPolicyInput) (*request.Request, *s3.GetBucketPolicyOutput)
GetBucketReplication(*s3.GetBucketReplicationInput) (*s3.GetBucketReplicationOutput, error)
GetBucketReplicationWithContext(aws.Context, *s3.GetBucketReplicationInput, ...request.Option) (*s3.GetBucketReplicationOutput, error)
GetBucketReplicationRequest(*s3.GetBucketReplicationInput) (*request.Request, *s3.GetBucketReplicationOutput)
GetBucketRequestPayment(*s3.GetBucketRequestPaymentInput) (*s3.GetBucketRequestPaymentOutput, error)
GetBucketRequestPaymentWithContext(aws.Context, *s3.GetBucketRequestPaymentInput, ...request.Option) (*s3.GetBucketRequestPaymentOutput, error)
GetBucketRequestPaymentRequest(*s3.GetBucketRequestPaymentInput) (*request.Request, *s3.GetBucketRequestPaymentOutput)
GetBucketTagging(*s3.GetBucketTaggingInput) (*s3.GetBucketTaggingOutput, error)
GetBucketTaggingWithContext(aws.Context, *s3.GetBucketTaggingInput, ...request.Option) (*s3.GetBucketTaggingOutput, error)
GetBucketTaggingRequest(*s3.GetBucketTaggingInput) (*request.Request, *s3.GetBucketTaggingOutput)
GetBucketVersioning(*s3.GetBucketVersioningInput) (*s3.GetBucketVersioningOutput, error)
GetBucketVersioningWithContext(aws.Context, *s3.GetBucketVersioningInput, ...request.Option) (*s3.GetBucketVersioningOutput, error)
GetBucketVersioningRequest(*s3.GetBucketVersioningInput) (*request.Request, *s3.GetBucketVersioningOutput)
GetBucketWebsite(*s3.GetBucketWebsiteInput) (*s3.GetBucketWebsiteOutput, error)
GetBucketWebsiteWithContext(aws.Context, *s3.GetBucketWebsiteInput, ...request.Option) (*s3.GetBucketWebsiteOutput, error)
GetBucketWebsiteRequest(*s3.GetBucketWebsiteInput) (*request.Request, *s3.GetBucketWebsiteOutput)
GetObject(*s3.GetObjectInput) (*s3.GetObjectOutput, error)
GetObjectWithContext(aws.Context, *s3.GetObjectInput, ...request.Option) (*s3.GetObjectOutput, error)
GetObjectRequest(*s3.GetObjectInput) (*request.Request, *s3.GetObjectOutput)
GetObjectAcl(*s3.GetObjectAclInput) (*s3.GetObjectAclOutput, error)
GetObjectAclWithContext(aws.Context, *s3.GetObjectAclInput, ...request.Option) (*s3.GetObjectAclOutput, error)
GetObjectAclRequest(*s3.GetObjectAclInput) (*request.Request, *s3.GetObjectAclOutput)
GetObjectTagging(*s3.GetObjectTaggingInput) (*s3.GetObjectTaggingOutput, error)
GetObjectTaggingWithContext(aws.Context, *s3.GetObjectTaggingInput, ...request.Option) (*s3.GetObjectTaggingOutput, error)
GetObjectTaggingRequest(*s3.GetObjectTaggingInput) (*request.Request, *s3.GetObjectTaggingOutput)
GetObjectTorrent(*s3.GetObjectTorrentInput) (*s3.GetObjectTorrentOutput, error)
GetObjectTorrentWithContext(aws.Context, *s3.GetObjectTorrentInput, ...request.Option) (*s3.GetObjectTorrentOutput, error)
GetObjectTorrentRequest(*s3.GetObjectTorrentInput) (*request.Request, *s3.GetObjectTorrentOutput)
HeadBucket(*s3.HeadBucketInput) (*s3.HeadBucketOutput, error)
HeadBucketWithContext(aws.Context, *s3.HeadBucketInput, ...request.Option) (*s3.HeadBucketOutput, error)
HeadBucketRequest(*s3.HeadBucketInput) (*request.Request, *s3.HeadBucketOutput)
HeadObject(*s3.HeadObjectInput) (*s3.HeadObjectOutput, error)
HeadObjectWithContext(aws.Context, *s3.HeadObjectInput, ...request.Option) (*s3.HeadObjectOutput, error)
HeadObjectRequest(*s3.HeadObjectInput) (*request.Request, *s3.HeadObjectOutput)
ListBucketAnalyticsConfigurations(*s3.ListBucketAnalyticsConfigurationsInput) (*s3.ListBucketAnalyticsConfigurationsOutput, error)
ListBucketAnalyticsConfigurationsWithContext(aws.Context, *s3.ListBucketAnalyticsConfigurationsInput, ...request.Option) (*s3.ListBucketAnalyticsConfigurationsOutput, error)
ListBucketAnalyticsConfigurationsRequest(*s3.ListBucketAnalyticsConfigurationsInput) (*request.Request, *s3.ListBucketAnalyticsConfigurationsOutput)
ListBucketInventoryConfigurations(*s3.ListBucketInventoryConfigurationsInput) (*s3.ListBucketInventoryConfigurationsOutput, error)
ListBucketInventoryConfigurationsWithContext(aws.Context, *s3.ListBucketInventoryConfigurationsInput, ...request.Option) (*s3.ListBucketInventoryConfigurationsOutput, error)
ListBucketInventoryConfigurationsRequest(*s3.ListBucketInventoryConfigurationsInput) (*request.Request, *s3.ListBucketInventoryConfigurationsOutput)
ListBucketMetricsConfigurations(*s3.ListBucketMetricsConfigurationsInput) (*s3.ListBucketMetricsConfigurationsOutput, error)
ListBucketMetricsConfigurationsWithContext(aws.Context, *s3.ListBucketMetricsConfigurationsInput, ...request.Option) (*s3.ListBucketMetricsConfigurationsOutput, error)
ListBucketMetricsConfigurationsRequest(*s3.ListBucketMetricsConfigurationsInput) (*request.Request, *s3.ListBucketMetricsConfigurationsOutput)
ListBuckets(*s3.ListBucketsInput) (*s3.ListBucketsOutput, error)
ListBucketsWithContext(aws.Context, *s3.ListBucketsInput, ...request.Option) (*s3.ListBucketsOutput, error)
ListBucketsRequest(*s3.ListBucketsInput) (*request.Request, *s3.ListBucketsOutput)
ListMultipartUploads(*s3.ListMultipartUploadsInput) (*s3.ListMultipartUploadsOutput, error)
ListMultipartUploadsWithContext(aws.Context, *s3.ListMultipartUploadsInput, ...request.Option) (*s3.ListMultipartUploadsOutput, error)
ListMultipartUploadsRequest(*s3.ListMultipartUploadsInput) (*request.Request, *s3.ListMultipartUploadsOutput)
ListMultipartUploadsPages(*s3.ListMultipartUploadsInput, func(*s3.ListMultipartUploadsOutput, bool) bool) error
ListMultipartUploadsPagesWithContext(aws.Context, *s3.ListMultipartUploadsInput, func(*s3.ListMultipartUploadsOutput, bool) bool, ...request.Option) error
ListObjectVersions(*s3.ListObjectVersionsInput) (*s3.ListObjectVersionsOutput, error)
ListObjectVersionsWithContext(aws.Context, *s3.ListObjectVersionsInput, ...request.Option) (*s3.ListObjectVersionsOutput, error)
ListObjectVersionsRequest(*s3.ListObjectVersionsInput) (*request.Request, *s3.ListObjectVersionsOutput)
ListObjectVersionsPages(*s3.ListObjectVersionsInput, func(*s3.ListObjectVersionsOutput, bool) bool) error
ListObjectVersionsPagesWithContext(aws.Context, *s3.ListObjectVersionsInput, func(*s3.ListObjectVersionsOutput, bool) bool, ...request.Option) error
ListObjects(*s3.ListObjectsInput) (*s3.ListObjectsOutput, error)
ListObjectsWithContext(aws.Context, *s3.ListObjectsInput, ...request.Option) (*s3.ListObjectsOutput, error)
ListObjectsRequest(*s3.ListObjectsInput) (*request.Request, *s3.ListObjectsOutput)
ListObjectsPages(*s3.ListObjectsInput, func(*s3.ListObjectsOutput, bool) bool) error
ListObjectsPagesWithContext(aws.Context, *s3.ListObjectsInput, func(*s3.ListObjectsOutput, bool) bool, ...request.Option) error
ListObjectsV2(*s3.ListObjectsV2Input) (*s3.ListObjectsV2Output, error)
ListObjectsV2WithContext(aws.Context, *s3.ListObjectsV2Input, ...request.Option) (*s3.ListObjectsV2Output, error)
ListObjectsV2Request(*s3.ListObjectsV2Input) (*request.Request, *s3.ListObjectsV2Output)
ListObjectsV2Pages(*s3.ListObjectsV2Input, func(*s3.ListObjectsV2Output, bool) bool) error
ListObjectsV2PagesWithContext(aws.Context, *s3.ListObjectsV2Input, func(*s3.ListObjectsV2Output, bool) bool, ...request.Option) error
ListParts(*s3.ListPartsInput) (*s3.ListPartsOutput, error)
ListPartsWithContext(aws.Context, *s3.ListPartsInput, ...request.Option) (*s3.ListPartsOutput, error)
ListPartsRequest(*s3.ListPartsInput) (*request.Request, *s3.ListPartsOutput)
ListPartsPages(*s3.ListPartsInput, func(*s3.ListPartsOutput, bool) bool) error
ListPartsPagesWithContext(aws.Context, *s3.ListPartsInput, func(*s3.ListPartsOutput, bool) bool, ...request.Option) error
PutBucketAccelerateConfiguration(*s3.PutBucketAccelerateConfigurationInput) (*s3.PutBucketAccelerateConfigurationOutput, error)
PutBucketAccelerateConfigurationWithContext(aws.Context, *s3.PutBucketAccelerateConfigurationInput, ...request.Option) (*s3.PutBucketAccelerateConfigurationOutput, error)
PutBucketAccelerateConfigurationRequest(*s3.PutBucketAccelerateConfigurationInput) (*request.Request, *s3.PutBucketAccelerateConfigurationOutput)
PutBucketAcl(*s3.PutBucketAclInput) (*s3.PutBucketAclOutput, error)
PutBucketAclWithContext(aws.Context, *s3.PutBucketAclInput, ...request.Option) (*s3.PutBucketAclOutput, error)
PutBucketAclRequest(*s3.PutBucketAclInput) (*request.Request, *s3.PutBucketAclOutput)
PutBucketAnalyticsConfiguration(*s3.PutBucketAnalyticsConfigurationInput) (*s3.PutBucketAnalyticsConfigurationOutput, error)
PutBucketAnalyticsConfigurationWithContext(aws.Context, *s3.PutBucketAnalyticsConfigurationInput, ...request.Option) (*s3.PutBucketAnalyticsConfigurationOutput, error)
PutBucketAnalyticsConfigurationRequest(*s3.PutBucketAnalyticsConfigurationInput) (*request.Request, *s3.PutBucketAnalyticsConfigurationOutput)
PutBucketCors(*s3.PutBucketCorsInput) (*s3.PutBucketCorsOutput, error)
PutBucketCorsWithContext(aws.Context, *s3.PutBucketCorsInput, ...request.Option) (*s3.PutBucketCorsOutput, error)
PutBucketCorsRequest(*s3.PutBucketCorsInput) (*request.Request, *s3.PutBucketCorsOutput)
PutBucketEncryption(*s3.PutBucketEncryptionInput) (*s3.PutBucketEncryptionOutput, error)
PutBucketEncryptionWithContext(aws.Context, *s3.PutBucketEncryptionInput, ...request.Option) (*s3.PutBucketEncryptionOutput, error)
PutBucketEncryptionRequest(*s3.PutBucketEncryptionInput) (*request.Request, *s3.PutBucketEncryptionOutput)
PutBucketInventoryConfiguration(*s3.PutBucketInventoryConfigurationInput) (*s3.PutBucketInventoryConfigurationOutput, error)
PutBucketInventoryConfigurationWithContext(aws.Context, *s3.PutBucketInventoryConfigurationInput, ...request.Option) (*s3.PutBucketInventoryConfigurationOutput, error)
PutBucketInventoryConfigurationRequest(*s3.PutBucketInventoryConfigurationInput) (*request.Request, *s3.PutBucketInventoryConfigurationOutput)
PutBucketLifecycle(*s3.PutBucketLifecycleInput) (*s3.PutBucketLifecycleOutput, error)
PutBucketLifecycleWithContext(aws.Context, *s3.PutBucketLifecycleInput, ...request.Option) (*s3.PutBucketLifecycleOutput, error)
PutBucketLifecycleRequest(*s3.PutBucketLifecycleInput) (*request.Request, *s3.PutBucketLifecycleOutput)
PutBucketLifecycleConfiguration(*s3.PutBucketLifecycleConfigurationInput) (*s3.PutBucketLifecycleConfigurationOutput, error)
PutBucketLifecycleConfigurationWithContext(aws.Context, *s3.PutBucketLifecycleConfigurationInput, ...request.Option) (*s3.PutBucketLifecycleConfigurationOutput, error)
PutBucketLifecycleConfigurationRequest(*s3.PutBucketLifecycleConfigurationInput) (*request.Request, *s3.PutBucketLifecycleConfigurationOutput)
PutBucketLogging(*s3.PutBucketLoggingInput) (*s3.PutBucketLoggingOutput, error)
PutBucketLoggingWithContext(aws.Context, *s3.PutBucketLoggingInput, ...request.Option) (*s3.PutBucketLoggingOutput, error)
PutBucketLoggingRequest(*s3.PutBucketLoggingInput) (*request.Request, *s3.PutBucketLoggingOutput)
PutBucketMetricsConfiguration(*s3.PutBucketMetricsConfigurationInput) (*s3.PutBucketMetricsConfigurationOutput, error)
PutBucketMetricsConfigurationWithContext(aws.Context, *s3.PutBucketMetricsConfigurationInput, ...request.Option) (*s3.PutBucketMetricsConfigurationOutput, error)
PutBucketMetricsConfigurationRequest(*s3.PutBucketMetricsConfigurationInput) (*request.Request, *s3.PutBucketMetricsConfigurationOutput)
PutBucketNotification(*s3.PutBucketNotificationInput) (*s3.PutBucketNotificationOutput, error)
PutBucketNotificationWithContext(aws.Context, *s3.PutBucketNotificationInput, ...request.Option) (*s3.PutBucketNotificationOutput, error)
PutBucketNotificationRequest(*s3.PutBucketNotificationInput) (*request.Request, *s3.PutBucketNotificationOutput)
PutBucketNotificationConfiguration(*s3.PutBucketNotificationConfigurationInput) (*s3.PutBucketNotificationConfigurationOutput, error)
PutBucketNotificationConfigurationWithContext(aws.Context, *s3.PutBucketNotificationConfigurationInput, ...request.Option) (*s3.PutBucketNotificationConfigurationOutput, error)
PutBucketNotificationConfigurationRequest(*s3.PutBucketNotificationConfigurationInput) (*request.Request, *s3.PutBucketNotificationConfigurationOutput)
PutBucketPolicy(*s3.PutBucketPolicyInput) (*s3.PutBucketPolicyOutput, error)
PutBucketPolicyWithContext(aws.Context, *s3.PutBucketPolicyInput, ...request.Option) (*s3.PutBucketPolicyOutput, error)
PutBucketPolicyRequest(*s3.PutBucketPolicyInput) (*request.Request, *s3.PutBucketPolicyOutput)
PutBucketReplication(*s3.PutBucketReplicationInput) (*s3.PutBucketReplicationOutput, error)
PutBucketReplicationWithContext(aws.Context, *s3.PutBucketReplicationInput, ...request.Option) (*s3.PutBucketReplicationOutput, error)
PutBucketReplicationRequest(*s3.PutBucketReplicationInput) (*request.Request, *s3.PutBucketReplicationOutput)
PutBucketRequestPayment(*s3.PutBucketRequestPaymentInput) (*s3.PutBucketRequestPaymentOutput, error)
PutBucketRequestPaymentWithContext(aws.Context, *s3.PutBucketRequestPaymentInput, ...request.Option) (*s3.PutBucketRequestPaymentOutput, error)
PutBucketRequestPaymentRequest(*s3.PutBucketRequestPaymentInput) (*request.Request, *s3.PutBucketRequestPaymentOutput)
PutBucketTagging(*s3.PutBucketTaggingInput) (*s3.PutBucketTaggingOutput, error)
PutBucketTaggingWithContext(aws.Context, *s3.PutBucketTaggingInput, ...request.Option) (*s3.PutBucketTaggingOutput, error)
PutBucketTaggingRequest(*s3.PutBucketTaggingInput) (*request.Request, *s3.PutBucketTaggingOutput)
PutBucketVersioning(*s3.PutBucketVersioningInput) (*s3.PutBucketVersioningOutput, error)
PutBucketVersioningWithContext(aws.Context, *s3.PutBucketVersioningInput, ...request.Option) (*s3.PutBucketVersioningOutput, error)
PutBucketVersioningRequest(*s3.PutBucketVersioningInput) (*request.Request, *s3.PutBucketVersioningOutput)
PutBucketWebsite(*s3.PutBucketWebsiteInput) (*s3.PutBucketWebsiteOutput, error)
PutBucketWebsiteWithContext(aws.Context, *s3.PutBucketWebsiteInput, ...request.Option) (*s3.PutBucketWebsiteOutput, error)
PutBucketWebsiteRequest(*s3.PutBucketWebsiteInput) (*request.Request, *s3.PutBucketWebsiteOutput)
PutObject(*s3.PutObjectInput) (*s3.PutObjectOutput, error)
PutObjectWithContext(aws.Context, *s3.PutObjectInput, ...request.Option) (*s3.PutObjectOutput, error)
PutObjectRequest(*s3.PutObjectInput) (*request.Request, *s3.PutObjectOutput)
PutObjectAcl(*s3.PutObjectAclInput) (*s3.PutObjectAclOutput, error)
PutObjectAclWithContext(aws.Context, *s3.PutObjectAclInput, ...request.Option) (*s3.PutObjectAclOutput, error)
PutObjectAclRequest(*s3.PutObjectAclInput) (*request.Request, *s3.PutObjectAclOutput)
PutObjectTagging(*s3.PutObjectTaggingInput) (*s3.PutObjectTaggingOutput, error)
PutObjectTaggingWithContext(aws.Context, *s3.PutObjectTaggingInput, ...request.Option) (*s3.PutObjectTaggingOutput, error)
PutObjectTaggingRequest(*s3.PutObjectTaggingInput) (*request.Request, *s3.PutObjectTaggingOutput)
RestoreObject(*s3.RestoreObjectInput) (*s3.RestoreObjectOutput, error)
RestoreObjectWithContext(aws.Context, *s3.RestoreObjectInput, ...request.Option) (*s3.RestoreObjectOutput, error)
RestoreObjectRequest(*s3.RestoreObjectInput) (*request.Request, *s3.RestoreObjectOutput)
SelectObjectContent(*s3.SelectObjectContentInput) (*s3.SelectObjectContentOutput, error)
SelectObjectContentWithContext(aws.Context, *s3.SelectObjectContentInput, ...request.Option) (*s3.SelectObjectContentOutput, error)
SelectObjectContentRequest(*s3.SelectObjectContentInput) (*request.Request, *s3.SelectObjectContentOutput)
UploadPart(*s3.UploadPartInput) (*s3.UploadPartOutput, error)
UploadPartWithContext(aws.Context, *s3.UploadPartInput, ...request.Option) (*s3.UploadPartOutput, error)
UploadPartRequest(*s3.UploadPartInput) (*request.Request, *s3.UploadPartOutput)
UploadPartCopy(*s3.UploadPartCopyInput) (*s3.UploadPartCopyOutput, error)
UploadPartCopyWithContext(aws.Context, *s3.UploadPartCopyInput, ...request.Option) (*s3.UploadPartCopyOutput, error)
UploadPartCopyRequest(*s3.UploadPartCopyInput) (*request.Request, *s3.UploadPartCopyOutput)
WaitUntilBucketExists(*s3.HeadBucketInput) error
WaitUntilBucketExistsWithContext(aws.Context, *s3.HeadBucketInput, ...request.WaiterOption) error
WaitUntilBucketNotExists(*s3.HeadBucketInput) error
WaitUntilBucketNotExistsWithContext(aws.Context, *s3.HeadBucketInput, ...request.WaiterOption) error
WaitUntilObjectExists(*s3.HeadObjectInput) error
WaitUntilObjectExistsWithContext(aws.Context, *s3.HeadObjectInput, ...request.WaiterOption) error
WaitUntilObjectNotExists(*s3.HeadObjectInput) error
WaitUntilObjectNotExistsWithContext(aws.Context, *s3.HeadObjectInput, ...request.WaiterOption) error
}
var _ S3API = (*s3.S3)(nil)
+529
View File
@@ -0,0 +1,529 @@
package s3manager
import (
"bytes"
"fmt"
"io"
"github.com/aws/aws-sdk-go/aws"
"github.com/aws/aws-sdk-go/aws/awserr"
"github.com/aws/aws-sdk-go/aws/client"
"github.com/aws/aws-sdk-go/aws/request"
"github.com/aws/aws-sdk-go/service/s3"
"github.com/aws/aws-sdk-go/service/s3/s3iface"
)
const (
// DefaultBatchSize is the batch size we initialize when constructing a batch delete client.
// This value is used when calling DeleteObjects. This represents how many objects to delete
// per DeleteObjects call.
DefaultBatchSize = 100
)
// BatchError will contain the key and bucket of the object that failed to
// either upload or download.
type BatchError struct {
Errors Errors
code string
message string
}
// Errors is a typed alias for a slice of errors to satisfy the error
// interface.
type Errors []Error
func (errs Errors) Error() string {
buf := bytes.NewBuffer(nil)
for i, err := range errs {
buf.WriteString(err.Error())
if i+1 < len(errs) {
buf.WriteString("\n")
}
}
return buf.String()
}
// Error will contain the original error, bucket, and key of the operation that failed
// during batch operations.
type Error struct {
OrigErr error
Bucket *string
Key *string
}
func newError(err error, bucket, key *string) Error {
return Error{
err,
bucket,
key,
}
}
func (err *Error) Error() string {
origErr := ""
if err.OrigErr != nil {
origErr = ":\n" + err.OrigErr.Error()
}
return fmt.Sprintf("failed to perform batch operation on %q to %q%s",
aws.StringValue(err.Key),
aws.StringValue(err.Bucket),
origErr,
)
}
// NewBatchError will return a BatchError that satisfies the awserr.Error interface.
func NewBatchError(code, message string, err []Error) awserr.Error {
return &BatchError{
Errors: err,
code: code,
message: message,
}
}
// Code will return the code associated with the batch error.
func (err *BatchError) Code() string {
return err.code
}
// Message will return the message associated with the batch error.
func (err *BatchError) Message() string {
return err.message
}
func (err *BatchError) Error() string {
return awserr.SprintError(err.Code(), err.Message(), "", err.Errors)
}
// OrigErr will return the original error. Which, in this case, will always be nil
// for batched operations.
func (err *BatchError) OrigErr() error {
return err.Errors
}
// BatchDeleteIterator is an interface that uses the scanner pattern to
// iterate through what needs to be deleted.
type BatchDeleteIterator interface {
Next() bool
Err() error
DeleteObject() BatchDeleteObject
}
// DeleteListIterator is an alternative iterator for the BatchDelete client. This will
// iterate through a list of objects and delete the objects.
//
// Example:
// iter := &s3manager.DeleteListIterator{
// Client: svc,
// Input: &s3.ListObjectsInput{
// Bucket: aws.String("bucket"),
// MaxKeys: aws.Int64(5),
// },
// Paginator: request.Pagination{
// NewRequest: func() (*request.Request, error) {
// var inCpy *ListObjectsInput
// if input != nil {
// tmp := *input
// inCpy = &tmp
// }
// req, _ := c.ListObjectsRequest(inCpy)
// return req, nil
// },
// },
// }
//
// batcher := s3manager.NewBatchDeleteWithClient(svc)
// if err := batcher.Delete(aws.BackgroundContext(), iter); err != nil {
// return err
// }
type DeleteListIterator struct {
Bucket *string
Paginator request.Pagination
objects []*s3.Object
}
// NewDeleteListIterator will return a new DeleteListIterator.
func NewDeleteListIterator(svc s3iface.S3API, input *s3.ListObjectsInput, opts ...func(*DeleteListIterator)) BatchDeleteIterator {
iter := &DeleteListIterator{
Bucket: input.Bucket,
Paginator: request.Pagination{
NewRequest: func() (*request.Request, error) {
var inCpy *s3.ListObjectsInput
if input != nil {
tmp := *input
inCpy = &tmp
}
req, _ := svc.ListObjectsRequest(inCpy)
return req, nil
},
},
}
for _, opt := range opts {
opt(iter)
}
return iter
}
// Next will use the S3API client to iterate through a list of objects.
func (iter *DeleteListIterator) Next() bool {
if len(iter.objects) > 0 {
iter.objects = iter.objects[1:]
}
if len(iter.objects) == 0 && iter.Paginator.Next() {
iter.objects = iter.Paginator.Page().(*s3.ListObjectsOutput).Contents
}
return len(iter.objects) > 0
}
// Err will return the last known error from Next.
func (iter *DeleteListIterator) Err() error {
return iter.Paginator.Err()
}
// DeleteObject will return the current object to be deleted.
func (iter *DeleteListIterator) DeleteObject() BatchDeleteObject {
return BatchDeleteObject{
Object: &s3.DeleteObjectInput{
Bucket: iter.Bucket,
Key: iter.objects[0].Key,
},
}
}
// BatchDelete will use the s3 package's service client to perform a batch
// delete.
type BatchDelete struct {
Client s3iface.S3API
BatchSize int
}
// NewBatchDeleteWithClient will return a new delete client that can delete a batched amount of
// objects.
//
// Example:
// batcher := s3manager.NewBatchDeleteWithClient(client, size)
//
// objects := []BatchDeleteObject{
// {
// Object: &s3.DeleteObjectInput {
// Key: aws.String("key"),
// Bucket: aws.String("bucket"),
// },
// },
// }
//
// if err := batcher.Delete(aws.BackgroundContext(), &s3manager.DeleteObjectsIterator{
// Objects: objects,
// }); err != nil {
// return err
// }
func NewBatchDeleteWithClient(client s3iface.S3API, options ...func(*BatchDelete)) *BatchDelete {
svc := &BatchDelete{
Client: client,
BatchSize: DefaultBatchSize,
}
for _, opt := range options {
opt(svc)
}
return svc
}
// NewBatchDelete will return a new delete client that can delete a batched amount of
// objects.
//
// Example:
// batcher := s3manager.NewBatchDelete(sess, size)
//
// objects := []BatchDeleteObject{
// {
// Object: &s3.DeleteObjectInput {
// Key: aws.String("key"),
// Bucket: aws.String("bucket"),
// },
// },
// }
//
// if err := batcher.Delete(aws.BackgroundContext(), &s3manager.DeleteObjectsIterator{
// Objects: objects,
// }); err != nil {
// return err
// }
func NewBatchDelete(c client.ConfigProvider, options ...func(*BatchDelete)) *BatchDelete {
client := s3.New(c)
return NewBatchDeleteWithClient(client, options...)
}
// BatchDeleteObject is a wrapper object for calling the batch delete operation.
type BatchDeleteObject struct {
Object *s3.DeleteObjectInput
// After will run after each iteration during the batch process. This function will
// be executed whether or not the request was successful.
After func() error
}
// DeleteObjectsIterator is an interface that uses the scanner pattern to iterate
// through a series of objects to be deleted.
type DeleteObjectsIterator struct {
Objects []BatchDeleteObject
index int
inc bool
}
// Next will increment the default iterator's index and and ensure that there
// is another object to iterator to.
func (iter *DeleteObjectsIterator) Next() bool {
if iter.inc {
iter.index++
} else {
iter.inc = true
}
return iter.index < len(iter.Objects)
}
// Err will return an error. Since this is just used to satisfy the BatchDeleteIterator interface
// this will only return nil.
func (iter *DeleteObjectsIterator) Err() error {
return nil
}
// DeleteObject will return the BatchDeleteObject at the current batched index.
func (iter *DeleteObjectsIterator) DeleteObject() BatchDeleteObject {
object := iter.Objects[iter.index]
return object
}
// Delete will use the iterator to queue up objects that need to be deleted.
// Once the batch size is met, this will call the deleteBatch function.
func (d *BatchDelete) Delete(ctx aws.Context, iter BatchDeleteIterator) error {
var errs []Error
objects := []BatchDeleteObject{}
var input *s3.DeleteObjectsInput
for iter.Next() {
o := iter.DeleteObject()
if input == nil {
input = initDeleteObjectsInput(o.Object)
}
parity := hasParity(input, o)
if parity {
input.Delete.Objects = append(input.Delete.Objects, &s3.ObjectIdentifier{
Key: o.Object.Key,
VersionId: o.Object.VersionId,
})
objects = append(objects, o)
}
if len(input.Delete.Objects) == d.BatchSize || !parity {
if err := deleteBatch(ctx, d, input, objects); err != nil {
errs = append(errs, err...)
}
objects = objects[:0]
input = nil
if !parity {
objects = append(objects, o)
input = initDeleteObjectsInput(o.Object)
input.Delete.Objects = append(input.Delete.Objects, &s3.ObjectIdentifier{
Key: o.Object.Key,
VersionId: o.Object.VersionId,
})
}
}
}
// iter.Next() could return false (above) plus populate iter.Err()
if iter.Err() != nil {
errs = append(errs, newError(iter.Err(), nil, nil))
}
if input != nil && len(input.Delete.Objects) > 0 {
if err := deleteBatch(ctx, d, input, objects); err != nil {
errs = append(errs, err...)
}
}
if len(errs) > 0 {
return NewBatchError("BatchedDeleteIncomplete", "some objects have failed to be deleted.", errs)
}
return nil
}
func initDeleteObjectsInput(o *s3.DeleteObjectInput) *s3.DeleteObjectsInput {
return &s3.DeleteObjectsInput{
Bucket: o.Bucket,
MFA: o.MFA,
RequestPayer: o.RequestPayer,
Delete: &s3.Delete{},
}
}
const (
// ErrDeleteBatchFailCode represents an error code which will be returned
// only when DeleteObjects.Errors has an error that does not contain a code.
ErrDeleteBatchFailCode = "DeleteBatchError"
errDefaultDeleteBatchMessage = "failed to delete"
)
// deleteBatch will delete a batch of items in the objects parameters.
func deleteBatch(ctx aws.Context, d *BatchDelete, input *s3.DeleteObjectsInput, objects []BatchDeleteObject) []Error {
errs := []Error{}
if result, err := d.Client.DeleteObjectsWithContext(ctx, input); err != nil {
for i := 0; i < len(input.Delete.Objects); i++ {
errs = append(errs, newError(err, input.Bucket, input.Delete.Objects[i].Key))
}
} else if len(result.Errors) > 0 {
for i := 0; i < len(result.Errors); i++ {
code := ErrDeleteBatchFailCode
msg := errDefaultDeleteBatchMessage
if result.Errors[i].Message != nil {
msg = *result.Errors[i].Message
}
if result.Errors[i].Code != nil {
code = *result.Errors[i].Code
}
errs = append(errs, newError(awserr.New(code, msg, err), input.Bucket, result.Errors[i].Key))
}
}
for _, object := range objects {
if object.After == nil {
continue
}
if err := object.After(); err != nil {
errs = append(errs, newError(err, object.Object.Bucket, object.Object.Key))
}
}
return errs
}
func hasParity(o1 *s3.DeleteObjectsInput, o2 BatchDeleteObject) bool {
if o1.Bucket != nil && o2.Object.Bucket != nil {
if *o1.Bucket != *o2.Object.Bucket {
return false
}
} else if o1.Bucket != o2.Object.Bucket {
return false
}
if o1.MFA != nil && o2.Object.MFA != nil {
if *o1.MFA != *o2.Object.MFA {
return false
}
} else if o1.MFA != o2.Object.MFA {
return false
}
if o1.RequestPayer != nil && o2.Object.RequestPayer != nil {
if *o1.RequestPayer != *o2.Object.RequestPayer {
return false
}
} else if o1.RequestPayer != o2.Object.RequestPayer {
return false
}
return true
}
// BatchDownloadIterator is an interface that uses the scanner pattern to iterate
// through a series of objects to be downloaded.
type BatchDownloadIterator interface {
Next() bool
Err() error
DownloadObject() BatchDownloadObject
}
// BatchDownloadObject contains all necessary information to run a batch operation once.
type BatchDownloadObject struct {
Object *s3.GetObjectInput
Writer io.WriterAt
// After will run after each iteration during the batch process. This function will
// be executed whether or not the request was successful.
After func() error
}
// DownloadObjectsIterator implements the BatchDownloadIterator interface and allows for batched
// download of objects.
type DownloadObjectsIterator struct {
Objects []BatchDownloadObject
index int
inc bool
}
// Next will increment the default iterator's index and and ensure that there
// is another object to iterator to.
func (batcher *DownloadObjectsIterator) Next() bool {
if batcher.inc {
batcher.index++
} else {
batcher.inc = true
}
return batcher.index < len(batcher.Objects)
}
// DownloadObject will return the BatchDownloadObject at the current batched index.
func (batcher *DownloadObjectsIterator) DownloadObject() BatchDownloadObject {
object := batcher.Objects[batcher.index]
return object
}
// Err will return an error. Since this is just used to satisfy the BatchDeleteIterator interface
// this will only return nil.
func (batcher *DownloadObjectsIterator) Err() error {
return nil
}
// BatchUploadIterator is an interface that uses the scanner pattern to
// iterate through what needs to be uploaded.
type BatchUploadIterator interface {
Next() bool
Err() error
UploadObject() BatchUploadObject
}
// UploadObjectsIterator implements the BatchUploadIterator interface and allows for batched
// upload of objects.
type UploadObjectsIterator struct {
Objects []BatchUploadObject
index int
inc bool
}
// Next will increment the default iterator's index and and ensure that there
// is another object to iterator to.
func (batcher *UploadObjectsIterator) Next() bool {
if batcher.inc {
batcher.index++
} else {
batcher.inc = true
}
return batcher.index < len(batcher.Objects)
}
// Err will return an error. Since this is just used to satisfy the BatchUploadIterator interface
// this will only return nil.
func (batcher *UploadObjectsIterator) Err() error {
return nil
}
// UploadObject will return the BatchUploadObject at the current batched index.
func (batcher *UploadObjectsIterator) UploadObject() BatchUploadObject {
object := batcher.Objects[batcher.index]
return object
}
// BatchUploadObject contains all necessary information to run a batch operation once.
type BatchUploadObject struct {
Object *UploadInput
// After will run after each iteration during the batch process. This function will
// be executed whether or not the request was successful.
After func() error
}
+88
View File
@@ -0,0 +1,88 @@
package s3manager
import (
"github.com/aws/aws-sdk-go/aws"
"github.com/aws/aws-sdk-go/aws/client"
"github.com/aws/aws-sdk-go/aws/credentials"
"github.com/aws/aws-sdk-go/aws/request"
"github.com/aws/aws-sdk-go/service/s3"
"github.com/aws/aws-sdk-go/service/s3/s3iface"
)
// GetBucketRegion will attempt to get the region for a bucket using the
// regionHint to determine which AWS partition to perform the query on.
//
// The request will not be signed, and will not use your AWS credentials.
//
// A "NotFound" error code will be returned if the bucket does not exist in the
// AWS partition the regionHint belongs to. If the regionHint parameter is an
// empty string GetBucketRegion will fallback to the ConfigProvider's region
// config. If the regionHint is empty, and the ConfigProvider does not have a
// region value, an error will be returned..
//
// For example to get the region of a bucket which exists in "eu-central-1"
// you could provide a region hint of "us-west-2".
//
// sess := session.Must(session.NewSession())
//
// bucket := "my-bucket"
// region, err := s3manager.GetBucketRegion(ctx, sess, bucket, "us-west-2")
// if err != nil {
// if aerr, ok := err.(awserr.Error); ok && aerr.Code() == "NotFound" {
// fmt.Fprintf(os.Stderr, "unable to find bucket %s's region not found\n", bucket)
// }
// return err
// }
// fmt.Printf("Bucket %s is in %s region\n", bucket, region)
//
func GetBucketRegion(ctx aws.Context, c client.ConfigProvider, bucket, regionHint string, opts ...request.Option) (string, error) {
var cfg aws.Config
if len(regionHint) != 0 {
cfg.Region = aws.String(regionHint)
}
svc := s3.New(c, &cfg)
return GetBucketRegionWithClient(ctx, svc, bucket, opts...)
}
const bucketRegionHeader = "X-Amz-Bucket-Region"
// GetBucketRegionWithClient is the same as GetBucketRegion with the exception
// that it takes a S3 service client instead of a Session. The regionHint is
// derived from the region the S3 service client was created in.
//
// See GetBucketRegion for more information.
func GetBucketRegionWithClient(ctx aws.Context, svc s3iface.S3API, bucket string, opts ...request.Option) (string, error) {
req, _ := svc.HeadBucketRequest(&s3.HeadBucketInput{
Bucket: aws.String(bucket),
})
req.Config.S3ForcePathStyle = aws.Bool(true)
req.Config.Credentials = credentials.AnonymousCredentials
req.SetContext(ctx)
// Disable HTTP redirects to prevent an invalid 301 from eating the response
// because Go's HTTP client will fail, and drop the response if an 301 is
// received without a location header. S3 will return a 301 without the
// location header for HeadObject API calls.
req.DisableFollowRedirects = true
var bucketRegion string
req.Handlers.Send.PushBack(func(r *request.Request) {
bucketRegion = r.HTTPResponse.Header.Get(bucketRegionHeader)
if len(bucketRegion) == 0 {
return
}
r.HTTPResponse.StatusCode = 200
r.HTTPResponse.Status = "OK"
r.Error = nil
})
req.ApplyOptions(opts...)
if err := req.Send(); err != nil {
return "", err
}
bucketRegion = s3.NormalizeBucketLocation(bucketRegion)
return bucketRegion, nil
}
+3
View File
@@ -0,0 +1,3 @@
// Package s3manager provides utilities to upload and download objects from
// S3 concurrently. Helpful for when working with large objects.
package s3manager
+555
View File
@@ -0,0 +1,555 @@
package s3manager
import (
"fmt"
"io"
"net/http"
"strconv"
"strings"
"sync"
"github.com/aws/aws-sdk-go/aws"
"github.com/aws/aws-sdk-go/aws/awserr"
"github.com/aws/aws-sdk-go/aws/awsutil"
"github.com/aws/aws-sdk-go/aws/client"
"github.com/aws/aws-sdk-go/aws/request"
"github.com/aws/aws-sdk-go/service/s3"
"github.com/aws/aws-sdk-go/service/s3/s3iface"
)
// DefaultDownloadPartSize is the default range of bytes to get at a time when
// using Download().
const DefaultDownloadPartSize = 1024 * 1024 * 5
// DefaultDownloadConcurrency is the default number of goroutines to spin up
// when using Download().
const DefaultDownloadConcurrency = 5
// The Downloader structure that calls Download(). It is safe to call Download()
// on this structure for multiple objects and across concurrent goroutines.
// Mutating the Downloader's properties is not safe to be done concurrently.
type Downloader struct {
// The buffer size (in bytes) to use when buffering data into chunks and
// sending them as parts to S3. The minimum allowed part size is 5MB, and
// if this value is set to zero, the DefaultDownloadPartSize value will be used.
//
// PartSize is ignored if the Range input parameter is provided.
PartSize int64
// The number of goroutines to spin up in parallel when sending parts.
// If this is set to zero, the DefaultDownloadConcurrency value will be used.
//
// Concurrency of 1 will download the parts sequentially.
//
// Concurrency is ignored if the Range input parameter is provided.
Concurrency int
// An S3 client to use when performing downloads.
S3 s3iface.S3API
// List of request options that will be passed down to individual API
// operation requests made by the downloader.
RequestOptions []request.Option
}
// WithDownloaderRequestOptions appends to the Downloader's API request options.
func WithDownloaderRequestOptions(opts ...request.Option) func(*Downloader) {
return func(d *Downloader) {
d.RequestOptions = append(d.RequestOptions, opts...)
}
}
// NewDownloader creates a new Downloader instance to downloads objects from
// S3 in concurrent chunks. Pass in additional functional options to customize
// the downloader behavior. Requires a client.ConfigProvider in order to create
// a S3 service client. The session.Session satisfies the client.ConfigProvider
// interface.
//
// Example:
// // The session the S3 Downloader will use
// sess := session.Must(session.NewSession())
//
// // Create a downloader with the session and default options
// downloader := s3manager.NewDownloader(sess)
//
// // Create a downloader with the session and custom options
// downloader := s3manager.NewDownloader(sess, func(d *s3manager.Downloader) {
// d.PartSize = 64 * 1024 * 1024 // 64MB per part
// })
func NewDownloader(c client.ConfigProvider, options ...func(*Downloader)) *Downloader {
d := &Downloader{
S3: s3.New(c),
PartSize: DefaultDownloadPartSize,
Concurrency: DefaultDownloadConcurrency,
}
for _, option := range options {
option(d)
}
return d
}
// NewDownloaderWithClient creates a new Downloader instance to downloads
// objects from S3 in concurrent chunks. Pass in additional functional
// options to customize the downloader behavior. Requires a S3 service client
// to make S3 API calls.
//
// Example:
// // The session the S3 Downloader will use
// sess := session.Must(session.NewSession())
//
// // The S3 client the S3 Downloader will use
// s3Svc := s3.new(sess)
//
// // Create a downloader with the s3 client and default options
// downloader := s3manager.NewDownloaderWithClient(s3Svc)
//
// // Create a downloader with the s3 client and custom options
// downloader := s3manager.NewDownloaderWithClient(s3Svc, func(d *s3manager.Downloader) {
// d.PartSize = 64 * 1024 * 1024 // 64MB per part
// })
func NewDownloaderWithClient(svc s3iface.S3API, options ...func(*Downloader)) *Downloader {
d := &Downloader{
S3: svc,
PartSize: DefaultDownloadPartSize,
Concurrency: DefaultDownloadConcurrency,
}
for _, option := range options {
option(d)
}
return d
}
type maxRetrier interface {
MaxRetries() int
}
// Download downloads an object in S3 and writes the payload into w using
// concurrent GET requests.
//
// Additional functional options can be provided to configure the individual
// download. These options are copies of the Downloader instance Download is called from.
// Modifying the options will not impact the original Downloader instance.
//
// It is safe to call this method concurrently across goroutines.
//
// The w io.WriterAt can be satisfied by an os.File to do multipart concurrent
// downloads, or in memory []byte wrapper using aws.WriteAtBuffer.
//
// Specifying a Downloader.Concurrency of 1 will cause the Downloader to
// download the parts from S3 sequentially.
//
// If the GetObjectInput's Range value is provided that will cause the downloader
// to perform a single GetObjectInput request for that object's range. This will
// caused the part size, and concurrency configurations to be ignored.
func (d Downloader) Download(w io.WriterAt, input *s3.GetObjectInput, options ...func(*Downloader)) (n int64, err error) {
return d.DownloadWithContext(aws.BackgroundContext(), w, input, options...)
}
// DownloadWithContext downloads an object in S3 and writes the payload into w
// using concurrent GET requests.
//
// DownloadWithContext is the same as Download with the additional support for
// Context input parameters. The Context must not be nil. A nil Context will
// cause a panic. Use the Context to add deadlining, timeouts, etc. The
// DownloadWithContext may create sub-contexts for individual underlying
// requests.
//
// Additional functional options can be provided to configure the individual
// download. These options are copies of the Downloader instance Download is
// called from. Modifying the options will not impact the original Downloader
// instance. Use the WithDownloaderRequestOptions helper function to pass in request
// options that will be applied to all API operations made with this downloader.
//
// The w io.WriterAt can be satisfied by an os.File to do multipart concurrent
// downloads, or in memory []byte wrapper using aws.WriteAtBuffer.
//
// Specifying a Downloader.Concurrency of 1 will cause the Downloader to
// download the parts from S3 sequentially.
//
// It is safe to call this method concurrently across goroutines.
//
// If the GetObjectInput's Range value is provided that will cause the downloader
// to perform a single GetObjectInput request for that object's range. This will
// caused the part size, and concurrency configurations to be ignored.
func (d Downloader) DownloadWithContext(ctx aws.Context, w io.WriterAt, input *s3.GetObjectInput, options ...func(*Downloader)) (n int64, err error) {
impl := downloader{w: w, in: input, cfg: d, ctx: ctx}
for _, option := range options {
option(&impl.cfg)
}
impl.cfg.RequestOptions = append(impl.cfg.RequestOptions, request.WithAppendUserAgent("S3Manager"))
if s, ok := d.S3.(maxRetrier); ok {
impl.partBodyMaxRetries = s.MaxRetries()
}
impl.totalBytes = -1
if impl.cfg.Concurrency == 0 {
impl.cfg.Concurrency = DefaultDownloadConcurrency
}
if impl.cfg.PartSize == 0 {
impl.cfg.PartSize = DefaultDownloadPartSize
}
return impl.download()
}
// DownloadWithIterator will download a batched amount of objects in S3 and writes them
// to the io.WriterAt specificed in the iterator.
//
// Example:
// svc := s3manager.NewDownloader(session)
//
// fooFile, err := os.Open("/tmp/foo.file")
// if err != nil {
// return err
// }
//
// barFile, err := os.Open("/tmp/bar.file")
// if err != nil {
// return err
// }
//
// objects := []s3manager.BatchDownloadObject {
// {
// Object: &s3.GetObjectInput {
// Bucket: aws.String("bucket"),
// Key: aws.String("foo"),
// },
// Writer: fooFile,
// },
// {
// Object: &s3.GetObjectInput {
// Bucket: aws.String("bucket"),
// Key: aws.String("bar"),
// },
// Writer: barFile,
// },
// }
//
// iter := &s3manager.DownloadObjectsIterator{Objects: objects}
// if err := svc.DownloadWithIterator(aws.BackgroundContext(), iter); err != nil {
// return err
// }
func (d Downloader) DownloadWithIterator(ctx aws.Context, iter BatchDownloadIterator, opts ...func(*Downloader)) error {
var errs []Error
for iter.Next() {
object := iter.DownloadObject()
if _, err := d.DownloadWithContext(ctx, object.Writer, object.Object, opts...); err != nil {
errs = append(errs, newError(err, object.Object.Bucket, object.Object.Key))
}
if object.After == nil {
continue
}
if err := object.After(); err != nil {
errs = append(errs, newError(err, object.Object.Bucket, object.Object.Key))
}
}
if len(errs) > 0 {
return NewBatchError("BatchedDownloadIncomplete", "some objects have failed to download.", errs)
}
return nil
}
// downloader is the implementation structure used internally by Downloader.
type downloader struct {
ctx aws.Context
cfg Downloader
in *s3.GetObjectInput
w io.WriterAt
wg sync.WaitGroup
m sync.Mutex
pos int64
totalBytes int64
written int64
err error
partBodyMaxRetries int
}
// download performs the implementation of the object download across ranged
// GETs.
func (d *downloader) download() (n int64, err error) {
// If range is specified fall back to single download of that range
// this enables the functionality of ranged gets with the downloader but
// at the cost of no multipart downloads.
if rng := aws.StringValue(d.in.Range); len(rng) > 0 {
d.downloadRange(rng)
return d.written, d.err
}
// Spin off first worker to check additional header information
d.getChunk()
if total := d.getTotalBytes(); total >= 0 {
// Spin up workers
ch := make(chan dlchunk, d.cfg.Concurrency)
for i := 0; i < d.cfg.Concurrency; i++ {
d.wg.Add(1)
go d.downloadPart(ch)
}
// Assign work
for d.getErr() == nil {
if d.pos >= total {
break // We're finished queuing chunks
}
// Queue the next range of bytes to read.
ch <- dlchunk{w: d.w, start: d.pos, size: d.cfg.PartSize}
d.pos += d.cfg.PartSize
}
// Wait for completion
close(ch)
d.wg.Wait()
} else {
// Checking if we read anything new
for d.err == nil {
d.getChunk()
}
// We expect a 416 error letting us know we are done downloading the
// total bytes. Since we do not know the content's length, this will
// keep grabbing chunks of data until the range of bytes specified in
// the request is out of range of the content. Once, this happens, a
// 416 should occur.
e, ok := d.err.(awserr.RequestFailure)
if ok && e.StatusCode() == http.StatusRequestedRangeNotSatisfiable {
d.err = nil
}
}
// Return error
return d.written, d.err
}
// downloadPart is an individual goroutine worker reading from the ch channel
// and performing a GetObject request on the data with a given byte range.
//
// If this is the first worker, this operation also resolves the total number
// of bytes to be read so that the worker manager knows when it is finished.
func (d *downloader) downloadPart(ch chan dlchunk) {
defer d.wg.Done()
for {
chunk, ok := <-ch
if !ok {
break
}
if d.getErr() != nil {
// Drain the channel if there is an error, to prevent deadlocking
// of download producer.
continue
}
if err := d.downloadChunk(chunk); err != nil {
d.setErr(err)
}
}
}
// getChunk grabs a chunk of data from the body.
// Not thread safe. Should only used when grabbing data on a single thread.
func (d *downloader) getChunk() {
if d.getErr() != nil {
return
}
chunk := dlchunk{w: d.w, start: d.pos, size: d.cfg.PartSize}
d.pos += d.cfg.PartSize
if err := d.downloadChunk(chunk); err != nil {
d.setErr(err)
}
}
// downloadRange downloads an Object given the passed in Byte-Range value.
// The chunk used down download the range will be configured for that range.
func (d *downloader) downloadRange(rng string) {
if d.getErr() != nil {
return
}
chunk := dlchunk{w: d.w, start: d.pos}
// Ranges specified will short circuit the multipart download
chunk.withRange = rng
if err := d.downloadChunk(chunk); err != nil {
d.setErr(err)
}
// Update the position based on the amount of data received.
d.pos = d.written
}
// downloadChunk downloads the chunk from s3
func (d *downloader) downloadChunk(chunk dlchunk) error {
in := &s3.GetObjectInput{}
awsutil.Copy(in, d.in)
// Get the next byte range of data
in.Range = aws.String(chunk.ByteRange())
var n int64
var err error
for retry := 0; retry <= d.partBodyMaxRetries; retry++ {
var resp *s3.GetObjectOutput
resp, err = d.cfg.S3.GetObjectWithContext(d.ctx, in, d.cfg.RequestOptions...)
if err != nil {
return err
}
d.setTotalBytes(resp) // Set total if not yet set.
n, err = io.Copy(&chunk, resp.Body)
resp.Body.Close()
if err == nil {
break
}
chunk.cur = 0
logMessage(d.cfg.S3, aws.LogDebugWithRequestRetries,
fmt.Sprintf("DEBUG: object part body download interrupted %s, err, %v, retrying attempt %d",
aws.StringValue(in.Key), err, retry))
}
d.incrWritten(n)
return err
}
func logMessage(svc s3iface.S3API, level aws.LogLevelType, msg string) {
s, ok := svc.(*s3.S3)
if !ok {
return
}
if s.Config.Logger == nil {
return
}
if s.Config.LogLevel.Matches(level) {
s.Config.Logger.Log(msg)
}
}
// getTotalBytes is a thread-safe getter for retrieving the total byte status.
func (d *downloader) getTotalBytes() int64 {
d.m.Lock()
defer d.m.Unlock()
return d.totalBytes
}
// setTotalBytes is a thread-safe setter for setting the total byte status.
// Will extract the object's total bytes from the Content-Range if the file
// will be chunked, or Content-Length. Content-Length is used when the response
// does not include a Content-Range. Meaning the object was not chunked. This
// occurs when the full file fits within the PartSize directive.
func (d *downloader) setTotalBytes(resp *s3.GetObjectOutput) {
d.m.Lock()
defer d.m.Unlock()
if d.totalBytes >= 0 {
return
}
if resp.ContentRange == nil {
// ContentRange is nil when the full file contents is provided, and
// is not chunked. Use ContentLength instead.
if resp.ContentLength != nil {
d.totalBytes = *resp.ContentLength
return
}
} else {
parts := strings.Split(*resp.ContentRange, "/")
total := int64(-1)
var err error
// Checking for whether or not a numbered total exists
// If one does not exist, we will assume the total to be -1, undefined,
// and sequentially download each chunk until hitting a 416 error
totalStr := parts[len(parts)-1]
if totalStr != "*" {
total, err = strconv.ParseInt(totalStr, 10, 64)
if err != nil {
d.err = err
return
}
}
d.totalBytes = total
}
}
func (d *downloader) incrWritten(n int64) {
d.m.Lock()
defer d.m.Unlock()
d.written += n
}
// getErr is a thread-safe getter for the error object
func (d *downloader) getErr() error {
d.m.Lock()
defer d.m.Unlock()
return d.err
}
// setErr is a thread-safe setter for the error object
func (d *downloader) setErr(e error) {
d.m.Lock()
defer d.m.Unlock()
d.err = e
}
// dlchunk represents a single chunk of data to write by the worker routine.
// This structure also implements an io.SectionReader style interface for
// io.WriterAt, effectively making it an io.SectionWriter (which does not
// exist).
type dlchunk struct {
w io.WriterAt
start int64
size int64
cur int64
// specifies the byte range the chunk should be downloaded with.
withRange string
}
// Write wraps io.WriterAt for the dlchunk, writing from the dlchunk's start
// position to its end (or EOF).
//
// If a range is specified on the dlchunk the size will be ignored when writing.
// as the total size may not of be known ahead of time.
func (c *dlchunk) Write(p []byte) (n int, err error) {
if c.cur >= c.size && len(c.withRange) == 0 {
return 0, io.EOF
}
n, err = c.w.WriteAt(p, c.start+c.cur)
c.cur += int64(n)
return
}
// ByteRange returns a HTTP Byte-Range header value that should be used by the
// client to request the chunk's range.
func (c *dlchunk) ByteRange() string {
if len(c.withRange) != 0 {
return c.withRange
}
return fmt.Sprintf("bytes=%d-%d", c.start, c.start+c.size-1)
}
+802
View File
@@ -0,0 +1,802 @@
package s3manager
import (
"bytes"
"fmt"
"io"
"sort"
"sync"
"time"
"github.com/aws/aws-sdk-go/aws"
"github.com/aws/aws-sdk-go/aws/awserr"
"github.com/aws/aws-sdk-go/aws/awsutil"
"github.com/aws/aws-sdk-go/aws/client"
"github.com/aws/aws-sdk-go/aws/request"
"github.com/aws/aws-sdk-go/service/s3"
"github.com/aws/aws-sdk-go/service/s3/s3iface"
)
// MaxUploadParts is the maximum allowed number of parts in a multi-part upload
// on Amazon S3.
const MaxUploadParts = 10000
// MinUploadPartSize is the minimum allowed part size when uploading a part to
// Amazon S3.
const MinUploadPartSize int64 = 1024 * 1024 * 5
// DefaultUploadPartSize is the default part size to buffer chunks of a
// payload into.
const DefaultUploadPartSize = MinUploadPartSize
// DefaultUploadConcurrency is the default number of goroutines to spin up when
// using Upload().
const DefaultUploadConcurrency = 5
// A MultiUploadFailure wraps a failed S3 multipart upload. An error returned
// will satisfy this interface when a multi part upload failed to upload all
// chucks to S3. In the case of a failure the UploadID is needed to operate on
// the chunks, if any, which were uploaded.
//
// Example:
//
// u := s3manager.NewUploader(opts)
// output, err := u.upload(input)
// if err != nil {
// if multierr, ok := err.(s3manager.MultiUploadFailure); ok {
// // Process error and its associated uploadID
// fmt.Println("Error:", multierr.Code(), multierr.Message(), multierr.UploadID())
// } else {
// // Process error generically
// fmt.Println("Error:", err.Error())
// }
// }
//
type MultiUploadFailure interface {
awserr.Error
// Returns the upload id for the S3 multipart upload that failed.
UploadID() string
}
// So that the Error interface type can be included as an anonymous field
// in the multiUploadError struct and not conflict with the error.Error() method.
type awsError awserr.Error
// A multiUploadError wraps the upload ID of a failed s3 multipart upload.
// Composed of BaseError for code, message, and original error
//
// Should be used for an error that occurred failing a S3 multipart upload,
// and a upload ID is available. If an uploadID is not available a more relevant
type multiUploadError struct {
awsError
// ID for multipart upload which failed.
uploadID string
}
// Error returns the string representation of the error.
//
// See apierr.BaseError ErrorWithExtra for output format
//
// Satisfies the error interface.
func (m multiUploadError) Error() string {
extra := fmt.Sprintf("upload id: %s", m.uploadID)
return awserr.SprintError(m.Code(), m.Message(), extra, m.OrigErr())
}
// String returns the string representation of the error.
// Alias for Error to satisfy the stringer interface.
func (m multiUploadError) String() string {
return m.Error()
}
// UploadID returns the id of the S3 upload which failed.
func (m multiUploadError) UploadID() string {
return m.uploadID
}
// UploadInput contains all input for upload requests to Amazon S3.
type UploadInput struct {
// The canned ACL to apply to the object.
ACL *string `location:"header" locationName:"x-amz-acl" type:"string"`
Bucket *string `location:"uri" locationName:"Bucket" type:"string" required:"true"`
// Specifies caching behavior along the request/reply chain.
CacheControl *string `location:"header" locationName:"Cache-Control" type:"string"`
// Specifies presentational information for the object.
ContentDisposition *string `location:"header" locationName:"Content-Disposition" type:"string"`
// Specifies what content encodings have been applied to the object and thus
// what decoding mechanisms must be applied to obtain the media-type referenced
// by the Content-Type header field.
ContentEncoding *string `location:"header" locationName:"Content-Encoding" type:"string"`
// The language the content is in.
ContentLanguage *string `location:"header" locationName:"Content-Language" type:"string"`
// The base64-encoded 128-bit MD5 digest of the part data.
ContentMD5 *string `location:"header" locationName:"Content-MD5" type:"string"`
// A standard MIME type describing the format of the object data.
ContentType *string `location:"header" locationName:"Content-Type" type:"string"`
// The date and time at which the object is no longer cacheable.
Expires *time.Time `location:"header" locationName:"Expires" type:"timestamp" timestampFormat:"rfc822"`
// Gives the grantee READ, READ_ACP, and WRITE_ACP permissions on the object.
GrantFullControl *string `location:"header" locationName:"x-amz-grant-full-control" type:"string"`
// Allows grantee to read the object data and its metadata.
GrantRead *string `location:"header" locationName:"x-amz-grant-read" type:"string"`
// Allows grantee to read the object ACL.
GrantReadACP *string `location:"header" locationName:"x-amz-grant-read-acp" type:"string"`
// Allows grantee to write the ACL for the applicable object.
GrantWriteACP *string `location:"header" locationName:"x-amz-grant-write-acp" type:"string"`
Key *string `location:"uri" locationName:"Key" type:"string" required:"true"`
// A map of metadata to store with the object in S3.
Metadata map[string]*string `location:"headers" locationName:"x-amz-meta-" type:"map"`
// Confirms that the requester knows that she or he will be charged for the
// request. Bucket owners need not specify this parameter in their requests.
// Documentation on downloading objects from requester pays buckets can be found
// at http://docs.aws.amazon.com/AmazonS3/latest/dev/ObjectsinRequesterPaysBuckets.html
RequestPayer *string `location:"header" locationName:"x-amz-request-payer" type:"string"`
// Specifies the algorithm to use to when encrypting the object (e.g., AES256,
// aws:kms).
SSECustomerAlgorithm *string `location:"header" locationName:"x-amz-server-side-encryption-customer-algorithm" type:"string"`
// Specifies the customer-provided encryption key for Amazon S3 to use in encrypting
// data. This value is used to store the object and then it is discarded; Amazon
// does not store the encryption key. The key must be appropriate for use with
// the algorithm specified in the x-amz-server-side​-encryption​-customer-algorithm
// header.
SSECustomerKey *string `location:"header" locationName:"x-amz-server-side-encryption-customer-key" type:"string"`
// Specifies the 128-bit MD5 digest of the encryption key according to RFC 1321.
// Amazon S3 uses this header for a message integrity check to ensure the encryption
// key was transmitted without error.
SSECustomerKeyMD5 *string `location:"header" locationName:"x-amz-server-side-encryption-customer-key-MD5" type:"string"`
// Specifies the AWS KMS key ID to use for object encryption. All GET and PUT
// requests for an object protected by AWS KMS will fail if not made via SSL
// or using SigV4. Documentation on configuring any of the officially supported
// AWS SDKs and CLI can be found at http://docs.aws.amazon.com/AmazonS3/latest/dev/UsingAWSSDK.html#specify-signature-version
SSEKMSKeyId *string `location:"header" locationName:"x-amz-server-side-encryption-aws-kms-key-id" type:"string"`
// The Server-side encryption algorithm used when storing this object in S3
// (e.g., AES256, aws:kms).
ServerSideEncryption *string `location:"header" locationName:"x-amz-server-side-encryption" type:"string"`
// The type of storage to use for the object. Defaults to 'STANDARD'.
StorageClass *string `location:"header" locationName:"x-amz-storage-class" type:"string"`
// The tag-set for the object. The tag-set must be encoded as URL Query parameters
Tagging *string `location:"header" locationName:"x-amz-tagging" type:"string"`
// If the bucket is configured as a website, redirects requests for this object
// to another object in the same bucket or to an external URL. Amazon S3 stores
// the value of this header in the object metadata.
WebsiteRedirectLocation *string `location:"header" locationName:"x-amz-website-redirect-location" type:"string"`
// The readable body payload to send to S3.
Body io.Reader
}
// UploadOutput represents a response from the Upload() call.
type UploadOutput struct {
// The URL where the object was uploaded to.
Location string
// The version of the object that was uploaded. Will only be populated if
// the S3 Bucket is versioned. If the bucket is not versioned this field
// will not be set.
VersionID *string
// The ID for a multipart upload to S3. In the case of an error the error
// can be cast to the MultiUploadFailure interface to extract the upload ID.
UploadID string
}
// WithUploaderRequestOptions appends to the Uploader's API request options.
func WithUploaderRequestOptions(opts ...request.Option) func(*Uploader) {
return func(u *Uploader) {
u.RequestOptions = append(u.RequestOptions, opts...)
}
}
// The Uploader structure that calls Upload(). It is safe to call Upload()
// on this structure for multiple objects and across concurrent goroutines.
// Mutating the Uploader's properties is not safe to be done concurrently.
type Uploader struct {
// The buffer size (in bytes) to use when buffering data into chunks and
// sending them as parts to S3. The minimum allowed part size is 5MB, and
// if this value is set to zero, the DefaultUploadPartSize value will be used.
PartSize int64
// The number of goroutines to spin up in parallel per call to Upload when
// sending parts. If this is set to zero, the DefaultUploadConcurrency value
// will be used.
//
// The concurrency pool is not shared between calls to Upload.
Concurrency int
// Setting this value to true will cause the SDK to avoid calling
// AbortMultipartUpload on a failure, leaving all successfully uploaded
// parts on S3 for manual recovery.
//
// Note that storing parts of an incomplete multipart upload counts towards
// space usage on S3 and will add additional costs if not cleaned up.
LeavePartsOnError bool
// MaxUploadParts is the max number of parts which will be uploaded to S3.
// Will be used to calculate the partsize of the object to be uploaded.
// E.g: 5GB file, with MaxUploadParts set to 100, will upload the file
// as 100, 50MB parts.
// With a limited of s3.MaxUploadParts (10,000 parts).
//
// Defaults to package const's MaxUploadParts value.
MaxUploadParts int
// The client to use when uploading to S3.
S3 s3iface.S3API
// List of request options that will be passed down to individual API
// operation requests made by the uploader.
RequestOptions []request.Option
}
// NewUploader creates a new Uploader instance to upload objects to S3. Pass In
// additional functional options to customize the uploader's behavior. Requires a
// client.ConfigProvider in order to create a S3 service client. The session.Session
// satisfies the client.ConfigProvider interface.
//
// Example:
// // The session the S3 Uploader will use
// sess := session.Must(session.NewSession())
//
// // Create an uploader with the session and default options
// uploader := s3manager.NewUploader(sess)
//
// // Create an uploader with the session and custom options
// uploader := s3manager.NewUploader(session, func(u *s3manager.Uploader) {
// u.PartSize = 64 * 1024 * 1024 // 64MB per part
// })
func NewUploader(c client.ConfigProvider, options ...func(*Uploader)) *Uploader {
u := &Uploader{
S3: s3.New(c),
PartSize: DefaultUploadPartSize,
Concurrency: DefaultUploadConcurrency,
LeavePartsOnError: false,
MaxUploadParts: MaxUploadParts,
}
for _, option := range options {
option(u)
}
return u
}
// NewUploaderWithClient creates a new Uploader instance to upload objects to S3. Pass in
// additional functional options to customize the uploader's behavior. Requires
// a S3 service client to make S3 API calls.
//
// Example:
// // The session the S3 Uploader will use
// sess := session.Must(session.NewSession())
//
// // S3 service client the Upload manager will use.
// s3Svc := s3.New(sess)
//
// // Create an uploader with S3 client and default options
// uploader := s3manager.NewUploaderWithClient(s3Svc)
//
// // Create an uploader with S3 client and custom options
// uploader := s3manager.NewUploaderWithClient(s3Svc, func(u *s3manager.Uploader) {
// u.PartSize = 64 * 1024 * 1024 // 64MB per part
// })
func NewUploaderWithClient(svc s3iface.S3API, options ...func(*Uploader)) *Uploader {
u := &Uploader{
S3: svc,
PartSize: DefaultUploadPartSize,
Concurrency: DefaultUploadConcurrency,
LeavePartsOnError: false,
MaxUploadParts: MaxUploadParts,
}
for _, option := range options {
option(u)
}
return u
}
// Upload uploads an object to S3, intelligently buffering large files into
// smaller chunks and sending them in parallel across multiple goroutines. You
// can configure the buffer size and concurrency through the Uploader's parameters.
//
// Additional functional options can be provided to configure the individual
// upload. These options are copies of the Uploader instance Upload is called from.
// Modifying the options will not impact the original Uploader instance.
//
// Use the WithUploaderRequestOptions helper function to pass in request
// options that will be applied to all API operations made with this uploader.
//
// It is safe to call this method concurrently across goroutines.
//
// Example:
// // Upload input parameters
// upParams := &s3manager.UploadInput{
// Bucket: &bucketName,
// Key: &keyName,
// Body: file,
// }
//
// // Perform an upload.
// result, err := uploader.Upload(upParams)
//
// // Perform upload with options different than the those in the Uploader.
// result, err := uploader.Upload(upParams, func(u *s3manager.Uploader) {
// u.PartSize = 10 * 1024 * 1024 // 10MB part size
// u.LeavePartsOnError = true // Don't delete the parts if the upload fails.
// })
func (u Uploader) Upload(input *UploadInput, options ...func(*Uploader)) (*UploadOutput, error) {
return u.UploadWithContext(aws.BackgroundContext(), input, options...)
}
// UploadWithContext uploads an object to S3, intelligently buffering large
// files into smaller chunks and sending them in parallel across multiple
// goroutines. You can configure the buffer size and concurrency through the
// Uploader's parameters.
//
// UploadWithContext is the same as Upload with the additional support for
// Context input parameters. The Context must not be nil. A nil Context will
// cause a panic. Use the context to add deadlining, timeouts, etc. The
// UploadWithContext may create sub-contexts for individual underlying requests.
//
// Additional functional options can be provided to configure the individual
// upload. These options are copies of the Uploader instance Upload is called from.
// Modifying the options will not impact the original Uploader instance.
//
// Use the WithUploaderRequestOptions helper function to pass in request
// options that will be applied to all API operations made with this uploader.
//
// It is safe to call this method concurrently across goroutines.
func (u Uploader) UploadWithContext(ctx aws.Context, input *UploadInput, opts ...func(*Uploader)) (*UploadOutput, error) {
i := uploader{in: input, cfg: u, ctx: ctx}
for _, opt := range opts {
opt(&i.cfg)
}
i.cfg.RequestOptions = append(i.cfg.RequestOptions, request.WithAppendUserAgent("S3Manager"))
return i.upload()
}
// UploadWithIterator will upload a batched amount of objects to S3. This operation uses
// the iterator pattern to know which object to upload next. Since this is an interface this
// allows for custom defined functionality.
//
// Example:
// svc:= s3manager.NewUploader(sess)
//
// objects := []BatchUploadObject{
// {
// Object: &s3manager.UploadInput {
// Key: aws.String("key"),
// Bucket: aws.String("bucket"),
// },
// },
// }
//
// iter := &s3manager.UploadObjectsIterator{Objects: objects}
// if err := svc.UploadWithIterator(aws.BackgroundContext(), iter); err != nil {
// return err
// }
func (u Uploader) UploadWithIterator(ctx aws.Context, iter BatchUploadIterator, opts ...func(*Uploader)) error {
var errs []Error
for iter.Next() {
object := iter.UploadObject()
if _, err := u.UploadWithContext(ctx, object.Object, opts...); err != nil {
s3Err := Error{
OrigErr: err,
Bucket: object.Object.Bucket,
Key: object.Object.Key,
}
errs = append(errs, s3Err)
}
if object.After == nil {
continue
}
if err := object.After(); err != nil {
s3Err := Error{
OrigErr: err,
Bucket: object.Object.Bucket,
Key: object.Object.Key,
}
errs = append(errs, s3Err)
}
}
if len(errs) > 0 {
return NewBatchError("BatchedUploadIncomplete", "some objects have failed to upload.", errs)
}
return nil
}
// internal structure to manage an upload to S3.
type uploader struct {
ctx aws.Context
cfg Uploader
in *UploadInput
readerPos int64 // current reader position
totalSize int64 // set to -1 if the size is not known
bufferPool sync.Pool
}
// internal logic for deciding whether to upload a single part or use a
// multipart upload.
func (u *uploader) upload() (*UploadOutput, error) {
u.init()
if u.cfg.PartSize < MinUploadPartSize {
msg := fmt.Sprintf("part size must be at least %d bytes", MinUploadPartSize)
return nil, awserr.New("ConfigError", msg, nil)
}
// Do one read to determine if we have more than one part
reader, _, part, err := u.nextReader()
if err == io.EOF { // single part
return u.singlePart(reader)
} else if err != nil {
return nil, awserr.New("ReadRequestBody", "read upload data failed", err)
}
mu := multiuploader{uploader: u}
return mu.upload(reader, part)
}
// init will initialize all default options.
func (u *uploader) init() {
if u.cfg.Concurrency == 0 {
u.cfg.Concurrency = DefaultUploadConcurrency
}
if u.cfg.PartSize == 0 {
u.cfg.PartSize = DefaultUploadPartSize
}
if u.cfg.MaxUploadParts == 0 {
u.cfg.MaxUploadParts = MaxUploadParts
}
u.bufferPool = sync.Pool{
New: func() interface{} { return make([]byte, u.cfg.PartSize) },
}
// Try to get the total size for some optimizations
u.initSize()
}
// initSize tries to detect the total stream size, setting u.totalSize. If
// the size is not known, totalSize is set to -1.
func (u *uploader) initSize() {
u.totalSize = -1
switch r := u.in.Body.(type) {
case io.Seeker:
n, err := aws.SeekerLen(r)
if err != nil {
return
}
u.totalSize = n
// Try to adjust partSize if it is too small and account for
// integer division truncation.
if u.totalSize/u.cfg.PartSize >= int64(u.cfg.MaxUploadParts) {
// Add one to the part size to account for remainders
// during the size calculation. e.g odd number of bytes.
u.cfg.PartSize = (u.totalSize / int64(u.cfg.MaxUploadParts)) + 1
}
}
}
// nextReader returns a seekable reader representing the next packet of data.
// This operation increases the shared u.readerPos counter, but note that it
// does not need to be wrapped in a mutex because nextReader is only called
// from the main thread.
func (u *uploader) nextReader() (io.ReadSeeker, int, []byte, error) {
type readerAtSeeker interface {
io.ReaderAt
io.ReadSeeker
}
switch r := u.in.Body.(type) {
case readerAtSeeker:
var err error
n := u.cfg.PartSize
if u.totalSize >= 0 {
bytesLeft := u.totalSize - u.readerPos
if bytesLeft <= u.cfg.PartSize {
err = io.EOF
n = bytesLeft
}
}
reader := io.NewSectionReader(r, u.readerPos, n)
u.readerPos += n
return reader, int(n), nil, err
default:
part := u.bufferPool.Get().([]byte)
n, err := readFillBuf(r, part)
u.readerPos += int64(n)
return bytes.NewReader(part[0:n]), n, part, err
}
}
func readFillBuf(r io.Reader, b []byte) (offset int, err error) {
for offset < len(b) && err == nil {
var n int
n, err = r.Read(b[offset:])
offset += n
}
return offset, err
}
// singlePart contains upload logic for uploading a single chunk via
// a regular PutObject request. Multipart requests require at least two
// parts, or at least 5MB of data.
func (u *uploader) singlePart(buf io.ReadSeeker) (*UploadOutput, error) {
params := &s3.PutObjectInput{}
awsutil.Copy(params, u.in)
params.Body = buf
// Need to use request form because URL generated in request is
// used in return.
req, out := u.cfg.S3.PutObjectRequest(params)
req.SetContext(u.ctx)
req.ApplyOptions(u.cfg.RequestOptions...)
if err := req.Send(); err != nil {
return nil, err
}
url := req.HTTPRequest.URL.String()
return &UploadOutput{
Location: url,
VersionID: out.VersionId,
}, nil
}
// internal structure to manage a specific multipart upload to S3.
type multiuploader struct {
*uploader
wg sync.WaitGroup
m sync.Mutex
err error
uploadID string
parts completedParts
}
// keeps track of a single chunk of data being sent to S3.
type chunk struct {
buf io.ReadSeeker
part []byte
num int64
}
// completedParts is a wrapper to make parts sortable by their part number,
// since S3 required this list to be sent in sorted order.
type completedParts []*s3.CompletedPart
func (a completedParts) Len() int { return len(a) }
func (a completedParts) Swap(i, j int) { a[i], a[j] = a[j], a[i] }
func (a completedParts) Less(i, j int) bool { return *a[i].PartNumber < *a[j].PartNumber }
// upload will perform a multipart upload using the firstBuf buffer containing
// the first chunk of data.
func (u *multiuploader) upload(firstBuf io.ReadSeeker, firstPart []byte) (*UploadOutput, error) {
params := &s3.CreateMultipartUploadInput{}
awsutil.Copy(params, u.in)
// Create the multipart
resp, err := u.cfg.S3.CreateMultipartUploadWithContext(u.ctx, params, u.cfg.RequestOptions...)
if err != nil {
return nil, err
}
u.uploadID = *resp.UploadId
// Create the workers
ch := make(chan chunk, u.cfg.Concurrency)
for i := 0; i < u.cfg.Concurrency; i++ {
u.wg.Add(1)
go u.readChunk(ch)
}
// Send part 1 to the workers
var num int64 = 1
ch <- chunk{buf: firstBuf, part: firstPart, num: num}
// Read and queue the rest of the parts
for u.geterr() == nil && err == nil {
num++
// This upload exceeded maximum number of supported parts, error now.
if num > int64(u.cfg.MaxUploadParts) || num > int64(MaxUploadParts) {
var msg string
if num > int64(u.cfg.MaxUploadParts) {
msg = fmt.Sprintf("exceeded total allowed configured MaxUploadParts (%d). Adjust PartSize to fit in this limit",
u.cfg.MaxUploadParts)
} else {
msg = fmt.Sprintf("exceeded total allowed S3 limit MaxUploadParts (%d). Adjust PartSize to fit in this limit",
MaxUploadParts)
}
u.seterr(awserr.New("TotalPartsExceeded", msg, nil))
break
}
var reader io.ReadSeeker
var nextChunkLen int
var part []byte
reader, nextChunkLen, part, err = u.nextReader()
if err != nil && err != io.EOF {
u.seterr(awserr.New(
"ReadRequestBody",
"read multipart upload data failed",
err))
break
}
if nextChunkLen == 0 {
// No need to upload empty part, if file was empty to start
// with empty single part would of been created and never
// started multipart upload.
break
}
ch <- chunk{buf: reader, part: part, num: num}
}
// Close the channel, wait for workers, and complete upload
close(ch)
u.wg.Wait()
complete := u.complete()
if err := u.geterr(); err != nil {
return nil, &multiUploadError{
awsError: awserr.New(
"MultipartUpload",
"upload multipart failed",
err),
uploadID: u.uploadID,
}
}
return &UploadOutput{
Location: aws.StringValue(complete.Location),
VersionID: complete.VersionId,
UploadID: u.uploadID,
}, nil
}
// readChunk runs in worker goroutines to pull chunks off of the ch channel
// and send() them as UploadPart requests.
func (u *multiuploader) readChunk(ch chan chunk) {
defer u.wg.Done()
for {
data, ok := <-ch
if !ok {
break
}
if u.geterr() == nil {
if err := u.send(data); err != nil {
u.seterr(err)
}
}
}
}
// send performs an UploadPart request and keeps track of the completed
// part information.
func (u *multiuploader) send(c chunk) error {
params := &s3.UploadPartInput{
Bucket: u.in.Bucket,
Key: u.in.Key,
Body: c.buf,
UploadId: &u.uploadID,
SSECustomerAlgorithm: u.in.SSECustomerAlgorithm,
SSECustomerKey: u.in.SSECustomerKey,
PartNumber: &c.num,
}
resp, err := u.cfg.S3.UploadPartWithContext(u.ctx, params, u.cfg.RequestOptions...)
// put the byte array back into the pool to conserve memory
u.bufferPool.Put(c.part)
if err != nil {
return err
}
n := c.num
completed := &s3.CompletedPart{ETag: resp.ETag, PartNumber: &n}
u.m.Lock()
u.parts = append(u.parts, completed)
u.m.Unlock()
return nil
}
// geterr is a thread-safe getter for the error object
func (u *multiuploader) geterr() error {
u.m.Lock()
defer u.m.Unlock()
return u.err
}
// seterr is a thread-safe setter for the error object
func (u *multiuploader) seterr(e error) {
u.m.Lock()
defer u.m.Unlock()
u.err = e
}
// fail will abort the multipart unless LeavePartsOnError is set to true.
func (u *multiuploader) fail() {
if u.cfg.LeavePartsOnError {
return
}
params := &s3.AbortMultipartUploadInput{
Bucket: u.in.Bucket,
Key: u.in.Key,
UploadId: &u.uploadID,
}
_, err := u.cfg.S3.AbortMultipartUploadWithContext(u.ctx, params, u.cfg.RequestOptions...)
if err != nil {
logMessage(u.cfg.S3, aws.LogDebug, fmt.Sprintf("failed to abort multipart upload, %v", err))
}
}
// complete successfully completes a multipart upload and returns the response.
func (u *multiuploader) complete() *s3.CompleteMultipartUploadOutput {
if u.geterr() != nil {
u.fail()
return nil
}
// Parts must be sorted in PartNumber order.
sort.Sort(u.parts)
params := &s3.CompleteMultipartUploadInput{
Bucket: u.in.Bucket,
Key: u.in.Key,
UploadId: &u.uploadID,
MultipartUpload: &s3.CompletedMultipartUpload{Parts: u.parts},
}
resp, err := u.cfg.S3.CompleteMultipartUploadWithContext(u.ctx, params, u.cfg.RequestOptions...)
if err != nil {
u.seterr(err)
u.fail()
}
return resp
}